diff --git a/cmd/gateway.go b/cmd/gateway.go index 61cb56e5..9ad41ca8 100644 --- a/cmd/gateway.go +++ b/cmd/gateway.go @@ -383,7 +383,7 @@ func runGateway() { httpapi.InitGatewayNoAuthFallbackAllowed(config.GatewayNoAuthFallbackAllowed(cfg.Gateway)) exportTokenStore := httpapi.InitExportTokenStore() defer exportTokenStore.Stop() - agentsH, skillsH, tracesH, mcpH, channelInstancesH, providersH, builtinToolsH, pendingMessagesH, teamEventsH, secureCLIH, secureCLIGrantH, mcpUserCredsH := wireHTTP(pgStores, cfg.Agents.Defaults.Workspace, dataDir, bundledSkillsDir, msgBus, toolsReg, providerRegistry, modelReg, permPE.IsOwner, gatewayAddr, mcpToolLister, usageCapSvc, cfg.Skills) + agentsH, skillsH, tracesH, mcpH, channelInstancesH, providersH, builtinToolsH, pendingMessagesH, teamEventsH, secureCLIH, secureCLIGrantH, mcpUserCredsH := wireHTTP(pgStores, cfg.Agents.Defaults.Workspace, dataDir, bundledSkillsDir, msgBus, toolsReg, providerRegistry, modelReg, permPE.IsOwner, gatewayAddr, mcpToolLister, usageCapSvc, cfg, cfg.Skills) // Wire dependencies for system prompt preview parity. if agentsH != nil { diff --git a/cmd/gateway_http_handlers.go b/cmd/gateway_http_handlers.go index 7d7cb532..2f255e3f 100644 --- a/cmd/gateway_http_handlers.go +++ b/cmd/gateway_http_handlers.go @@ -11,7 +11,7 @@ import ( ) // wireHTTP creates HTTP handlers (agents + skills + traces + MCP + channel instances + providers + builtin tools + pending messages). -func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir string, msgBus *bus.MessageBus, toolsReg *tools.Registry, providerReg *providers.Registry, modelReg providers.ModelRegistry, isOwner func(string) bool, gatewayAddr string, mcpToolLister httpapi.MCPToolLister, usageCapSvc *usagecaps.Service, skillUploadConfig config.SkillsConfig) (*httpapi.AgentsHandler, *httpapi.SkillsHandler, *httpapi.TracesHandler, *httpapi.MCPHandler, *httpapi.ChannelInstancesHandler, *httpapi.ProvidersHandler, *httpapi.BuiltinToolsHandler, *httpapi.PendingMessagesHandler, *httpapi.TeamEventsHandler, *httpapi.SecureCLIHandler, *httpapi.SecureCLIGrantHandler, *httpapi.MCPUserCredentialsHandler) { +func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir string, msgBus *bus.MessageBus, toolsReg *tools.Registry, providerReg *providers.Registry, modelReg providers.ModelRegistry, isOwner func(string) bool, gatewayAddr string, mcpToolLister httpapi.MCPToolLister, usageCapSvc *usagecaps.Service, appCfg *config.Config, skillUploadConfig config.SkillsConfig) (*httpapi.AgentsHandler, *httpapi.SkillsHandler, *httpapi.TracesHandler, *httpapi.MCPHandler, *httpapi.ChannelInstancesHandler, *httpapi.ProvidersHandler, *httpapi.BuiltinToolsHandler, *httpapi.PendingMessagesHandler, *httpapi.TeamEventsHandler, *httpapi.SecureCLIHandler, *httpapi.SecureCLIGrantHandler, *httpapi.MCPUserCredentialsHandler) { var agentsH *httpapi.AgentsHandler var skillsH *httpapi.SkillsHandler var tracesH *httpapi.TracesHandler @@ -68,6 +68,11 @@ func wireHTTP(stores *store.Stores, defaultWorkspace, dataDir, bundledSkillsDir providersH = httpapi.NewProvidersHandler(stores.Providers, stores.ConfigSecrets, providerReg, gatewayAddr) providersH.SetMessageBus(msgBus) providersH.SetUsageCapService(usageCapSvc) + if appCfg != nil { + providersH.SetShellDenyGroupsSource(func() map[string]bool { + return appCfg.ShellDenyGroupsSnapshot() + }) + } if modelReg != nil { providersH.SetModelRegistry(modelReg) } diff --git a/cmd/gateway_lifecycle.go b/cmd/gateway_lifecycle.go index 631bdd4e..d39b4c86 100644 --- a/cmd/gateway_lifecycle.go +++ b/cmd/gateway_lifecycle.go @@ -88,6 +88,13 @@ func (d *gatewayDeps) runLifecycle( // Reload global shell deny-group toggles on config changes via pub/sub // so /config edits apply without a process restart. subscribeShellDenyGroupsReload(d.msgBus, d.toolsReg) + var providerStore store.ProviderStore + var mcpStore store.MCPServerStore + if d.pgStores != nil { + providerStore = d.pgStores.Providers + mcpStore = d.pgStores.MCP + } + subscribeProviderShellDenyGroupsReload(d.msgBus, d.providerRegistry, providerStore, mcpStore) // Reload TTS providers on config changes via pub/sub. d.msgBus.Subscribe("tts-config-reload", func(evt bus.Event) { diff --git a/cmd/gateway_lifecycle_shell_deny_groups.go b/cmd/gateway_lifecycle_shell_deny_groups.go index 46e0e0df..d204e7a8 100644 --- a/cmd/gateway_lifecycle_shell_deny_groups.go +++ b/cmd/gateway_lifecycle_shell_deny_groups.go @@ -1,10 +1,13 @@ package cmd import ( + "context" "log/slog" "github.com/nextlevelbuilder/goclaw/internal/bus" "github.com/nextlevelbuilder/goclaw/internal/config" + "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" ) @@ -21,6 +24,7 @@ func subscribeShellDenyGroupsReload(msgBus *bus.MessageBus, toolsReg *tools.Regi if !ok { return } + snapshot := updatedCfg.Clone() execTool, ok := toolsReg.Get("exec") if !ok { return @@ -29,11 +33,58 @@ func subscribeShellDenyGroupsReload(msgBus *bus.MessageBus, toolsReg *tools.Regi if !ok { return } - et.SetGlobalShellDenyGroups(updatedCfg.Tools.ShellDenyGroups) - et.SetCommandKeywordAllowlist(updatedCfg.Tools.CommandKeywordAllowlist) + et.SetGlobalShellDenyGroups(snapshot.Tools.ShellDenyGroups) + et.SetCommandKeywordAllowlist(snapshot.Tools.CommandKeywordAllowlist) slog.Info("shell deny groups reloaded via pub/sub", - "groups", len(updatedCfg.Tools.ShellDenyGroups), - "command_keyword_allowlist_rules", len(updatedCfg.Tools.CommandKeywordAllowlist), + "groups", len(snapshot.Tools.ShellDenyGroups), + "command_keyword_allowlist_rules", len(snapshot.Tools.CommandKeywordAllowlist), ) }) } + +func subscribeProviderShellDenyGroupsReload(msgBus *bus.MessageBus, providerReg *providers.Registry, provStore store.ProviderStore, mcpStore store.MCPServerStore) { + if msgBus == nil || providerReg == nil { + return + } + msgBus.Subscribe("shell-deny-provider-policy-reload", func(evt bus.Event) { + if evt.Name != bus.TopicConfigChanged { + return + } + updatedCfg, ok := evt.Payload.(*config.Config) + if !ok { + return + } + reloadShellDenyProviderPolicies(providerReg, provStore, mcpStore, updatedCfg) + }) +} + +func reloadShellDenyProviderPolicies(providerReg *providers.Registry, provStore store.ProviderStore, mcpStore store.MCPServerStore, cfg *config.Config) { + if providerReg == nil || cfg == nil { + return + } + snapshot := cfg.Clone() + registerClaudeCLIFromConfig(providerReg, snapshot) + if snapshot.Providers.ACP.Binary != "" { + registerACPFromConfig(providerReg, snapshot.Providers.ACP, snapshot.ShellDenyGroupsSnapshot()) + } + if provStore == nil { + return + } + dbProviders, err := provStore.ListAllProviders(context.Background()) + if err != nil { + slog.Warn("shell deny provider policy reload: failed to load providers from DB", "error", err) + return + } + gatewayAddr := loopbackAddr(snapshot.Gateway.Host, snapshot.Gateway.Port) + for _, p := range dbProviders { + if !p.Enabled { + continue + } + switch p.ProviderType { + case store.ProviderClaudeCLI: + registerClaudeCLIFromDB(providerReg, p, gatewayAddr, snapshot.Gateway.Token, mcpStore, snapshot) + case store.ProviderACP: + registerACPFromDB(providerReg, p, snapshot.ShellDenyGroupsSnapshot()) + } + } +} diff --git a/cmd/gateway_lifecycle_shell_deny_groups_test.go b/cmd/gateway_lifecycle_shell_deny_groups_test.go index e66be5c2..0bd385e0 100644 --- a/cmd/gateway_lifecycle_shell_deny_groups_test.go +++ b/cmd/gateway_lifecycle_shell_deny_groups_test.go @@ -2,10 +2,18 @@ package cmd import ( "context" + "errors" + "os" + "path/filepath" + "regexp" "testing" + "github.com/google/uuid" + "github.com/nextlevelbuilder/goclaw/internal/bus" "github.com/nextlevelbuilder/goclaw/internal/config" + "github.com/nextlevelbuilder/goclaw/internal/providers" + "github.com/nextlevelbuilder/goclaw/internal/store" "github.com/nextlevelbuilder/goclaw/internal/tools" ) @@ -119,3 +127,153 @@ func TestShellDenyGroupsConfigReload_IgnoresWrongPayload(t *testing.T) { t.Fatalf("expected wrong-payload event to be ignored; package_install changed to %v", v) } } + +func TestShellDenyGroupsConfigReload_ReplacesConfigClaudeCLIProvider(t *testing.T) { + msgBus := bus.New() + defer msgBus.Unsubscribe("shell-deny-provider-policy-reload") + + providerReg := providers.NewRegistry(store.TenantIDFromContext) + defer providerReg.Close() + + initial := config.Default() + initial.Providers.ClaudeCLI.CLIPath = "claude" + initial.Tools.ShellDenyGroups = map[string]bool{"package_install": true} + reloadShellDenyProviderPolicies(providerReg, nil, nil, initial) + + before, err := providerReg.Get(context.Background(), "claude-cli") + if err != nil { + t.Fatal(err) + } + beforeCLI, ok := before.(*providers.ClaudeCLIProvider) + if !ok { + t.Fatalf("expected ClaudeCLI provider, got %T", before) + } + + updated := config.Default() + updated.Providers.ClaudeCLI.CLIPath = "claude" + updated.Tools.ShellDenyGroups = map[string]bool{"package_install": false} + subscribeProviderShellDenyGroupsReload(msgBus, providerReg, nil, nil) + msgBus.Broadcast(bus.Event{Name: bus.TopicConfigChanged, Payload: updated}) + + after, err := providerReg.Get(context.Background(), "claude-cli") + if err != nil { + t.Fatal(err) + } + afterCLI, ok := after.(*providers.ClaudeCLIProvider) + if !ok { + t.Fatalf("expected ClaudeCLI provider after reload, got %T", after) + } + if beforeCLI == afterCLI { + t.Fatal("expected config change to replace Claude CLI provider runtime") + } +} + +func TestShellDenyGroupsConfigReload_ReplacesDBClaudeCLIProvider(t *testing.T) { + providerReg := providers.NewRegistry(store.TenantIDFromContext) + defer providerReg.Close() + + tenantID := uuid.New() + binary := writeTestExecutable(t) + provStore := &shellDenyGroupsProviderStore{providers: []store.LLMProviderData{ + { + TenantID: tenantID, + Name: "tenant-claude", + ProviderType: store.ProviderClaudeCLI, + APIBase: binary, + Enabled: true, + }, + }} + + initial := config.Default() + initial.Tools.ShellDenyGroups = map[string]bool{"package_install": true} + reloadShellDenyProviderPolicies(providerReg, provStore, nil, initial) + + before, err := providerReg.GetForTenant(tenantID, "tenant-claude") + if err != nil { + t.Fatal(err) + } + beforeCLI, ok := before.(*providers.ClaudeCLIProvider) + if !ok { + t.Fatalf("expected ClaudeCLI provider, got %T", before) + } + + updated := config.Default() + updated.Tools.ShellDenyGroups = map[string]bool{"package_install": false} + reloadShellDenyProviderPolicies(providerReg, provStore, nil, updated) + + after, err := providerReg.GetForTenant(tenantID, "tenant-claude") + if err != nil { + t.Fatal(err) + } + afterCLI, ok := after.(*providers.ClaudeCLIProvider) + if !ok { + t.Fatalf("expected ClaudeCLI provider after reload, got %T", after) + } + if beforeCLI == afterCLI { + t.Fatal("expected config change to replace DB Claude CLI provider runtime") + } +} + +func TestConfiguredShellDenyPatternsDropsDisabledPackageInstall(t *testing.T) { + cfg := config.Default() + cfg.Tools.ShellDenyGroups = map[string]bool{"package_install": false} + + patterns := configuredShellDenyPatterns(cfg) + + if matchesAny(patterns, "pip install requests") { + t.Fatal("package_install=false should remove package-install deny patterns") + } + if !matchesAny(patterns, "env") { + t.Fatal("unrelated default deny patterns should remain active") + } +} + +type shellDenyGroupsProviderStore struct { + providers []store.LLMProviderData +} + +func (s *shellDenyGroupsProviderStore) CreateProvider(context.Context, *store.LLMProviderData) error { + return errors.New("not implemented") +} + +func (s *shellDenyGroupsProviderStore) GetProvider(context.Context, uuid.UUID) (*store.LLMProviderData, error) { + return nil, errors.New("not implemented") +} + +func (s *shellDenyGroupsProviderStore) GetProviderByName(context.Context, string) (*store.LLMProviderData, error) { + return nil, errors.New("not implemented") +} + +func (s *shellDenyGroupsProviderStore) ListProviders(context.Context) ([]store.LLMProviderData, error) { + return s.providers, nil +} + +func (s *shellDenyGroupsProviderStore) ListAllProviders(context.Context) ([]store.LLMProviderData, error) { + return s.providers, nil +} + +func (s *shellDenyGroupsProviderStore) UpdateProvider(context.Context, uuid.UUID, map[string]any) error { + return errors.New("not implemented") +} + +func (s *shellDenyGroupsProviderStore) DeleteProvider(context.Context, uuid.UUID) error { + return errors.New("not implemented") +} + +func writeTestExecutable(t *testing.T) string { + t.Helper() + binary := filepath.Join(t.TempDir(), "claude-test") + if err := os.WriteFile(binary, []byte("#!/bin/sh\nexit 0\n"), 0755); err != nil { + t.Fatal(err) + } + return binary +} + +func matchesAny(patterns []*regexp.Regexp, command string) bool { + for _, pattern := range patterns { + if pattern.MatchString(command) { + return true + } + } + return false +} diff --git a/cmd/gateway_managed.go b/cmd/gateway_managed.go index 5352c5bf..87809779 100644 --- a/cmd/gateway_managed.go +++ b/cmd/gateway_managed.go @@ -701,7 +701,7 @@ func wireExtras( } providerReg.UnregisterForTenant(tenantID, p.Name) if p.Enabled { - registerACPFromDB(providerReg, *p) + registerACPFromDB(providerReg, *p, configuredShellDenyGroups(appCfg)) } }) diff --git a/cmd/gateway_providers.go b/cmd/gateway_providers.go index ba19ff98..2e4c5967 100644 --- a/cmd/gateway_providers.go +++ b/cmd/gateway_providers.go @@ -7,6 +7,7 @@ import ( "net" "os/exec" "path/filepath" + "regexp" "strconv" "time" @@ -195,33 +196,11 @@ func registerProviders(registry *providers.Registry, cfg *config.Config, modelRe } } - // Claude CLI provider (subscription-based, no API key needed) - if cfg.Providers.ClaudeCLI.CLIPath != "" { - cliPath := cfg.Providers.ClaudeCLI.CLIPath - var opts []providers.ClaudeCLIOption - if cfg.Providers.ClaudeCLI.Model != "" { - opts = append(opts, providers.WithClaudeCLIModel(cfg.Providers.ClaudeCLI.Model)) - } - if cfg.Providers.ClaudeCLI.BaseWorkDir != "" { - opts = append(opts, providers.WithClaudeCLIWorkDir(cfg.Providers.ClaudeCLI.BaseWorkDir)) - } - if cfg.Providers.ClaudeCLI.PermMode != "" { - opts = append(opts, providers.WithClaudeCLIPermMode(cfg.Providers.ClaudeCLI.PermMode)) - } - // Build per-session MCP config: external MCP servers + GoClaw bridge - gatewayAddr := loopbackAddr(cfg.Gateway.Host, cfg.Gateway.Port) - mcpData := providers.BuildCLIMCPConfigData(cfg.Tools.McpServers, gatewayAddr, cfg.Gateway.Token) - opts = append(opts, providers.WithClaudeCLIMCPConfigData(mcpData)) - // Enable GoClaw security hooks (shell deny patterns, path restrictions) - opts = append(opts, providers.WithClaudeCLISecurityHooks( - cfg.Providers.ClaudeCLI.BaseWorkDir, true)) - registry.Register(providers.NewClaudeCLIProvider(cliPath, opts...)) - slog.Info("registered provider", "name", "claude-cli") - } + registerClaudeCLIFromConfig(registry, cfg) // ACP provider (config-based) — orchestrates any ACP-compatible agent binary if cfg.Providers.ACP.Binary != "" { - registerACPFromConfig(registry, cfg.Providers.ACP) + registerACPFromConfig(registry, cfg.Providers.ACP, configuredShellDenyGroups(cfg)) } } @@ -303,34 +282,12 @@ func registerProvidersFromDB(registry *providers.Registry, provStore store.Provi continue } if p.ProviderType == store.ProviderClaudeCLI { - cliPath := p.APIBase // reuse APIBase field for CLI path - if cliPath == "" { - cliPath = "claude" - } - // Validate: only accept "claude" or absolute path - if cliPath != "claude" && !filepath.IsAbs(cliPath) { - slog.Warn("security.claude_cli: invalid path from DB, using default", "path", cliPath) - cliPath = "claude" - } - if _, err := exec.LookPath(cliPath); err != nil { - slog.Warn("claude-cli: binary not found, skipping", "path", cliPath, "error", err) - continue - } - var cliOpts []providers.ClaudeCLIOption - cliOpts = append(cliOpts, providers.WithClaudeCLIName(p.Name)) - cliOpts = append(cliOpts, providers.WithClaudeCLISecurityHooks("", true)) - if gatewayAddr != "" { - mcpData := providers.BuildCLIMCPConfigData(nil, gatewayAddr, gatewayToken) - mcpData.AgentMCPLookup = buildMCPServerLookup(mcpStore) - cliOpts = append(cliOpts, providers.WithClaudeCLIMCPConfigData(mcpData)) - } - registry.RegisterForTenant(p.TenantID, providers.NewClaudeCLIProvider(cliPath, cliOpts...)) - slog.Info("registered provider from DB", "name", p.Name) + registerClaudeCLIFromDB(registry, p, gatewayAddr, gatewayToken, mcpStore, cfg) continue } // ACP provider — no API key needed (agents manage their own auth). if p.ProviderType == store.ProviderACP { - registerACPFromDB(registry, p) + registerACPFromDB(registry, p, configuredShellDenyGroups(cfg)) continue } // Local Ollama requires no API key — handle before the key guard (same pattern as ClaudeCLI). @@ -455,8 +412,59 @@ func registerProvidersFromDB(registry *providers.Registry, provStore store.Provi } } +func registerClaudeCLIFromConfig(registry *providers.Registry, cfg *config.Config) { + if cfg == nil || cfg.Providers.ClaudeCLI.CLIPath == "" { + return + } + cliPath := cfg.Providers.ClaudeCLI.CLIPath + var opts []providers.ClaudeCLIOption + if cfg.Providers.ClaudeCLI.Model != "" { + opts = append(opts, providers.WithClaudeCLIModel(cfg.Providers.ClaudeCLI.Model)) + } + if cfg.Providers.ClaudeCLI.BaseWorkDir != "" { + opts = append(opts, providers.WithClaudeCLIWorkDir(cfg.Providers.ClaudeCLI.BaseWorkDir)) + } + if cfg.Providers.ClaudeCLI.PermMode != "" { + opts = append(opts, providers.WithClaudeCLIPermMode(cfg.Providers.ClaudeCLI.PermMode)) + } + gatewayAddr := loopbackAddr(cfg.Gateway.Host, cfg.Gateway.Port) + mcpData := providers.BuildCLIMCPConfigData(cfg.Tools.McpServers, gatewayAddr, cfg.Gateway.Token) + opts = append(opts, providers.WithClaudeCLIMCPConfigData(mcpData)) + opts = append(opts, providers.WithClaudeCLISecurityHooks( + cfg.Providers.ClaudeCLI.BaseWorkDir, true, configuredShellDenyPatterns(cfg))) + registry.Register(providers.NewClaudeCLIProvider(cliPath, opts...)) + slog.Info("registered provider", "name", "claude-cli") +} + +func registerClaudeCLIFromDB(registry *providers.Registry, p store.LLMProviderData, gatewayAddr, gatewayToken string, mcpStore store.MCPServerStore, cfg *config.Config) bool { + cliPath := p.APIBase // reuse APIBase field for CLI path + if cliPath == "" { + cliPath = "claude" + } + // Validate: only accept "claude" or absolute path + if cliPath != "claude" && !filepath.IsAbs(cliPath) { + slog.Warn("security.claude_cli: invalid path from DB, using default", "path", cliPath) + cliPath = "claude" + } + if _, err := exec.LookPath(cliPath); err != nil { + slog.Warn("claude-cli: binary not found, skipping", "path", cliPath, "error", err) + return false + } + var cliOpts []providers.ClaudeCLIOption + cliOpts = append(cliOpts, providers.WithClaudeCLIName(p.Name)) + cliOpts = append(cliOpts, providers.WithClaudeCLISecurityHooks("", true, configuredShellDenyPatterns(cfg))) + if gatewayAddr != "" { + mcpData := providers.BuildCLIMCPConfigData(nil, gatewayAddr, gatewayToken) + mcpData.AgentMCPLookup = buildMCPServerLookup(mcpStore) + cliOpts = append(cliOpts, providers.WithClaudeCLIMCPConfigData(mcpData)) + } + registry.RegisterForTenant(p.TenantID, providers.NewClaudeCLIProvider(cliPath, cliOpts...)) + slog.Info("registered provider from DB", "name", p.Name) + return true +} + // registerACPFromConfig registers an ACP provider from config file settings. -func registerACPFromConfig(registry *providers.Registry, cfg config.ACPConfig) { +func registerACPFromConfig(registry *providers.Registry, cfg config.ACPConfig, shellDenyGroups map[string]bool) { if _, err := exec.LookPath(cfg.Binary); err != nil { slog.Warn("acp: binary not found, skipping", "binary", cfg.Binary, "error", err) return @@ -479,13 +487,13 @@ func registerACPFromConfig(registry *providers.Registry, cfg config.ACPConfig) { opts = append(opts, providers.WithACPPermMode(cfg.PermMode)) } registry.Register(providers.NewACPProvider( - cfg.Binary, cfg.Args, workDir, idleTTL, tools.DefaultDenyPatterns(), opts..., + cfg.Binary, cfg.Args, workDir, idleTTL, tools.ResolveDenyPatterns(shellDenyGroups), opts..., )) slog.Info("registered provider", "name", "acp", "binary", cfg.Binary) } // registerACPFromDB registers an ACP provider from a DB provider row. -func registerACPFromDB(registry *providers.Registry, p store.LLMProviderData) { +func registerACPFromDB(registry *providers.Registry, p store.LLMProviderData, shellDenyGroups map[string]bool) { binary := p.APIBase // repurpose api_base as binary path if binary == "" { slog.Warn("acp: no binary specified in DB provider", "name", p.Name) @@ -522,13 +530,24 @@ func registerACPFromDB(registry *providers.Registry, p store.LLMProviderData) { workDir = defaultACPWorkDir() } registry.RegisterForTenant(p.TenantID, providers.NewACPProvider( - binary, settings.Args, workDir, idleTTL, tools.DefaultDenyPatterns(), + binary, settings.Args, workDir, idleTTL, tools.ResolveDenyPatterns(shellDenyGroups), providers.WithACPName(p.Name), providers.WithACPModel(p.Name), )) slog.Info("registered provider from DB", "name", p.Name, "type", "acp") } +func configuredShellDenyGroups(cfg *config.Config) map[string]bool { + if cfg == nil { + return nil + } + return cfg.ShellDenyGroupsSnapshot() +} + +func configuredShellDenyPatterns(cfg *config.Config) []*regexp.Regexp { + return tools.ResolveDenyPatterns(configuredShellDenyGroups(cfg)) +} + // defaultACPWorkDir returns the default workspace directory for ACP agents. func defaultACPWorkDir() string { return filepath.Join(config.ResolvedDataDirFromEnv(), "acp-workspaces") diff --git a/docs/project-changelog.md b/docs/project-changelog.md index 7da561f2..93825551 100644 --- a/docs/project-changelog.md +++ b/docs/project-changelog.md @@ -6,6 +6,36 @@ Significant changes, features, and fixes in reverse chronological order. ## 2026-05-29 +### Shell security group disabled state persistence (issue #75) + +**Fixes** + +- Added regression coverage proving `config.patch` preserves explicit + `tools.shellDenyGroups` values set to `false` in memory, saved config JSON, + and follow-up `config.get` responses. +- Aligned Claude CLI and ACP provider runtime registration with the saved + shell deny-group config. Provider-side deny patterns now derive from + `tools.ResolveDenyPatterns(cfg.Tools.ShellDenyGroups)` instead of static + defaults, so disabled groups stay disabled after reload/provider registration. +- Re-registers existing Claude CLI/ACP provider runtimes on config changes so + settings-page saves affect runtime enforcement without restart. +- Uses locked shell deny-group snapshots before HTTP provider registration, so + provider creation cannot race config reload while reading the override map. +- Kept missing shell deny-group keys as inherited defaults; only explicit + `false` disables a group. + +**Tests** + +- `go test ./internal/gateway/methods -run 'TestConfigPatchPersists(InboundDebounceMs|ShellDenyGroupsFalse)' -count=1` +- `go test ./internal/providers -run 'TestGenerateHookScript' -count=1` +- `go test ./internal/providers/... -run 'ClaudeCLI|ACP|DenyPatterns|ToolBridge' -count=1` +- `go test -race ./cmd -run 'TestShellDenyGroupsConfigReload_.*|TestConfiguredShellDenyPatternsDropsDisabledPackageInstall' -count=1` +- `go vet ./...` +- `go build ./...` +- `go build -tags sqliteonly ./...` + +--- + ### Local-first document text extraction adapter (read_document privacy optimization) Adds optional local extraction pipeline to the `read_document` tool, allowing diff --git a/internal/config/config.go b/internal/config/config.go index 3f405911..f171e308 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -4,6 +4,7 @@ import ( "encoding/json" "fmt" "log/slog" + "maps" "strings" "sync" "time" @@ -132,7 +133,7 @@ type DatabaseConfig struct { type SkillsConfig struct { StorageDir string `json:"storage_dir,omitempty"` // directory for skill content (default: dataDir/skills-store/) MaxUploadSizeMB int `json:"max_upload_size_mb,omitempty"` // per-file ZIP upload limit - SlashCommands SkillSlashCommandConfig `json:"slash_commands,omitempty"` + SlashCommands SkillSlashCommandConfig `json:"slash_commands"` } // SkillSlashCommandConfig controls explicit slash-command skill activation. @@ -565,6 +566,36 @@ func (c *Config) ReplaceFrom(src *Config) { c.Bindings = src.Bindings } +// Clone returns a deep copy of the config while holding the read lock. +func (c *Config) Clone() *Config { + c.mu.RLock() + defer c.mu.RUnlock() + + data, err := json.Marshal(c) + if err != nil { + return &Config{} + } + cp := Default() + if err := json.Unmarshal(data, cp); err != nil { + return &Config{} + } + return cp +} + +// ShellDenyGroupsSnapshot returns a copy of the current global shell deny-group +// overrides. Callers can safely resolve patterns without racing config reloads. +func (c *Config) ShellDenyGroupsSnapshot() map[string]bool { + c.mu.RLock() + defer c.mu.RUnlock() + + if len(c.Tools.ShellDenyGroups) == 0 { + return nil + } + groups := make(map[string]bool, len(c.Tools.ShellDenyGroups)) + maps.Copy(groups, c.Tools.ShellDenyGroups) + return groups +} + // IdentityConfig defines agent persona / display identity. type IdentityConfig struct { Name string `json:"name,omitempty"` diff --git a/internal/gateway/methods/config.go b/internal/gateway/methods/config.go index 481181b3..d068d055 100644 --- a/internal/gateway/methods/config.go +++ b/internal/gateway/methods/config.go @@ -178,19 +178,8 @@ func (m *ConfigMethods) handlePatch(ctx context.Context, client *gateway.Client, return } - // Merge strategy: serialize current -> deserialize patch on top -> save - currentJSON, err := json.Marshal(m.cfg) - if err != nil { - client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInternal, i18n.T(locale, i18n.MsgInternalError, "failed to serialize current config"))) - return - } - - // Start from current config as base - merged := config.Default() - if err := json.Unmarshal(currentJSON, merged); err != nil { - client.SendResponse(protocol.NewErrorResponse(req.ID, protocol.ErrInternal, i18n.T(locale, i18n.MsgInternalError, "failed to clone config"))) - return - } + // Start from a locked current-config snapshot as base. + merged := m.cfg.Clone() // Apply patch on top if err := json5.Unmarshal([]byte(params.Raw), merged); err != nil { diff --git a/internal/gateway/methods/config_patch_test.go b/internal/gateway/methods/config_patch_test.go index cd86e722..b080e877 100644 --- a/internal/gateway/methods/config_patch_test.go +++ b/internal/gateway/methods/config_patch_test.go @@ -57,6 +57,80 @@ func TestConfigPatchPersistsInboundDebounceMs(t *testing.T) { } } +func TestConfigPatchPersistsShellDenyGroupsFalse(t *testing.T) { + t.Parallel() + + cfg := config.Default() + cfgPath := filepath.Join(t.TempDir(), "config.json") + methods := NewConfigMethods(cfg, cfgPath, nil, nil) + client, responses := gateway.NewCapturingTestClient(permissions.RoleOwner, store.MasterTenantID, "owner", 1) + params, err := json.Marshal(map[string]string{ + "raw": `{"tools":{"shellDenyGroups":{"package_install":false}}}`, + }) + if err != nil { + t.Fatal(err) + } + + ctx := store.WithTenantID(context.Background(), store.MasterTenantID) + methods.handlePatch( + ctx, + client, + &protocol.RequestFrame{ + Type: protocol.FrameTypeRequest, + ID: "patch-shell-deny-groups", + Method: protocol.MethodConfigPatch, + Params: params, + }, + ) + + res := readConfigPatchResponse(t, responses) + if !res.OK { + t.Fatalf("config.patch failed: %#v", res.Error) + } + shellDenyGroups := cfg.ShellDenyGroupsSnapshot() + if v, ok := shellDenyGroups["package_install"]; !ok || v { + t.Fatalf("in-memory shellDenyGroups package_install = %v (ok=%v), want false", v, ok) + } + data, err := os.ReadFile(cfgPath) + if err != nil { + t.Fatal(err) + } + if !bytes.Contains(data, []byte(`"package_install": false`)) { + t.Fatalf("saved config missing package_install=false:\n%s", data) + } + + methods.handleGet( + ctx, + client, + &protocol.RequestFrame{ + Type: protocol.FrameTypeRequest, + ID: "get-shell-deny-groups", + Method: protocol.MethodConfigGet, + }, + ) + getRes := readConfigPatchResponse(t, responses) + if !getRes.OK { + t.Fatalf("config.get failed: %#v", getRes.Error) + } + var payload struct { + Config struct { + Tools struct { + ShellDenyGroups map[string]bool `json:"shellDenyGroups"` + } `json:"tools"` + } `json:"config"` + } + rawPayload, err := json.Marshal(getRes.Payload) + if err != nil { + t.Fatal(err) + } + if err := json.Unmarshal(rawPayload, &payload); err != nil { + t.Fatal(err) + } + if v, ok := payload.Config.Tools.ShellDenyGroups["package_install"]; !ok || v { + t.Fatalf("config.get shellDenyGroups package_install = %v (ok=%v), want false", v, ok) + } +} + func readConfigPatchResponse(t *testing.T, responses <-chan []byte) protocol.ResponseFrame { t.Helper() select { diff --git a/internal/http/providers.go b/internal/http/providers.go index 9127982f..8cc46d82 100644 --- a/internal/http/providers.go +++ b/internal/http/providers.go @@ -10,6 +10,7 @@ import ( "net/url" "os/exec" "path/filepath" + "regexp" "strings" "sync" @@ -22,6 +23,7 @@ import ( "github.com/nextlevelbuilder/goclaw/internal/permissions" "github.com/nextlevelbuilder/goclaw/internal/providers" "github.com/nextlevelbuilder/goclaw/internal/store" + "github.com/nextlevelbuilder/goclaw/internal/tools" usagecaps "github.com/nextlevelbuilder/goclaw/internal/usage/caps" "github.com/nextlevelbuilder/goclaw/pkg/protocol" ) @@ -33,6 +35,7 @@ type ProvidersHandler struct { providerReg *providers.Registry gatewayAddr string // for injecting MCP bridge into Claude CLI providers mcpLookup providers.MCPServerLookup // optional: resolves per-agent MCP servers + shellDenyGroups func() map[string]bool // optional: current global shell deny-group overrides apiBaseFallback func(providerType string) string // optional: config/env fallback for api_base cliMu sync.Mutex // serializes Claude CLI provider create to prevent duplicates msgBus *bus.MessageBus @@ -65,6 +68,12 @@ func (h *ProvidersHandler) SetMCPServerLookup(lookup providers.MCPServerLookup) h.mcpLookup = lookup } +// SetShellDenyGroupsSource sets the current global shell deny-group source for +// runtime provider registration. Must be called before serving requests. +func (h *ProvidersHandler) SetShellDenyGroupsSource(fn func() map[string]bool) { + h.shellDenyGroups = fn +} + // SetAPIBaseFallback sets a function that returns config/env api_base by provider type. // Used as fallback when DB providers have no api_base set. func (h *ProvidersHandler) SetAPIBaseFallback(fn func(providerType string) string) { @@ -91,6 +100,13 @@ func (h *ProvidersHandler) SetUsageCapService(s *usagecaps.Service) { h.usageCaps = s } +func (h *ProvidersHandler) currentShellDenyPatterns() []*regexp.Regexp { + if h.shellDenyGroups == nil { + return tools.DefaultDenyPatterns() + } + return tools.ResolveDenyPatterns(h.shellDenyGroups()) +} + // resolveAPIBase returns the provider's api_base, falling back to config/env if empty. // For Ollama/OllamaCloud providers, applies a safety-net normalization: if the stored // value is missing the /v1 suffix (pre-existing record before write-time normalization), @@ -217,7 +233,7 @@ func (h *ProvidersHandler) registerInMemory(p *store.LLMProviderData) providerRu } cliOpts := []providers.ClaudeCLIOption{ providers.WithClaudeCLIName(p.Name), - providers.WithClaudeCLISecurityHooks("", true), + providers.WithClaudeCLISecurityHooks("", true, h.currentShellDenyPatterns()), } if h.gatewayAddr != "" { mcpData := providers.BuildCLIMCPConfigData(nil, h.gatewayAddr, pkgGatewayToken) diff --git a/internal/providers/claude_cli.go b/internal/providers/claude_cli.go index e5e0dde7..2271f5d0 100644 --- a/internal/providers/claude_cli.go +++ b/internal/providers/claude_cli.go @@ -3,6 +3,7 @@ package providers import ( "log/slog" "os" + "regexp" "sync" ) @@ -42,17 +43,17 @@ const OptLocalKey = "local_key" // It acts as a thin proxy: CLI manages session history, tool execution, and context. // GoClaw only forwards the latest user message and streams back the response. type ClaudeCLIProvider struct { - name string // provider name (default: "claude-cli") - cliPath string // path to claude binary (default: "claude") - defaultModel string // default: "sonnet" - baseWorkDir string // base dir for agent workspaces - mcpConfigData *MCPConfigData // per-session MCP config data - permMode string // permission mode (default: "bypassPermissions") - hooksSettingsPath string // generated settings.json with security hooks (empty = no hooks) - hooksCleanup func() // cleanup function for hooks temp files - mu sync.Mutex // protects workdir creation - sessionMu sync.Map // key: string, value: *sync.Mutex — per-session lock - mcpConfigDirs sync.Map // key: string (dir path), value: struct{} — tracks per-session MCP config dirs for cleanup + name string // provider name (default: "claude-cli") + cliPath string // path to claude binary (default: "claude") + defaultModel string // default: "sonnet" + baseWorkDir string // base dir for agent workspaces + mcpConfigData *MCPConfigData // per-session MCP config data + permMode string // permission mode (default: "bypassPermissions") + hooksSettingsPath string // generated settings.json with security hooks (empty = no hooks) + hooksCleanup func() // cleanup function for hooks temp files + mu sync.Mutex // protects workdir creation + sessionMu sync.Map // key: string, value: *sync.Mutex — per-session lock + mcpConfigDirs sync.Map // key: string (dir path), value: struct{} — tracks per-session MCP config dirs for cleanup } // ClaudeCLIOption configures the provider. @@ -105,9 +106,9 @@ func WithClaudeCLIPermMode(mode string) ClaudeCLIOption { // WithClaudeCLISecurityHooks enables GoClaw security hooks for CLI tool calls. // Generates a settings file with PreToolUse hooks that enforce shell deny patterns // and workspace path restrictions. -func WithClaudeCLISecurityHooks(workspace string, restrictToWorkspace bool) ClaudeCLIOption { +func WithClaudeCLISecurityHooks(workspace string, restrictToWorkspace bool, denyPatternSets ...[]*regexp.Regexp) ClaudeCLIOption { return func(p *ClaudeCLIProvider) { - settingsPath, cleanup, err := BuildCLIHooksConfig(workspace, restrictToWorkspace) + settingsPath, cleanup, err := BuildCLIHooksConfig(workspace, restrictToWorkspace, denyPatternSets...) if err != nil { slog.Warn("claude-cli: failed to build security hooks", "error", err) return @@ -136,7 +137,7 @@ func NewClaudeCLIProvider(cliPath string, opts ...ClaudeCLIOption) *ClaudeCLIPro return p } -func (p *ClaudeCLIProvider) Name() string { return p.name } +func (p *ClaudeCLIProvider) Name() string { return p.name } func (p *ClaudeCLIProvider) DefaultModel() string { return p.defaultModel } // Capabilities implements CapabilitiesAware for pipeline code-path selection. diff --git a/internal/providers/claude_cli_hooks.go b/internal/providers/claude_cli_hooks.go index 004c4192..f7d71936 100644 --- a/internal/providers/claude_cli_hooks.go +++ b/internal/providers/claude_cli_hooks.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "regexp" "strings" "github.com/google/uuid" @@ -13,7 +14,7 @@ import ( // BuildCLIHooksConfig generates a Claude CLI settings file with PreToolUse hooks // that enforce GoClaw's security policies (shell deny patterns, path restrictions). // Returns settings file path and a cleanup function. -func BuildCLIHooksConfig(workspace string, restrictToWorkspace bool) (string, func(), error) { +func BuildCLIHooksConfig(workspace string, restrictToWorkspace bool, denyPatternSets ...[]*regexp.Regexp) (string, func(), error) { tmpDir := filepath.Join(os.TempDir(), "goclaw-cli-hooks") if err := os.MkdirAll(tmpDir, 0755); err != nil { return "", nil, fmt.Errorf("create hooks dir: %w", err) @@ -22,7 +23,7 @@ func BuildCLIHooksConfig(workspace string, restrictToWorkspace bool) (string, fu id := uuid.New().String()[:8] // Write the hook script - hookScript := generateHookScript(workspace, restrictToWorkspace) + hookScript := generateHookScript(workspace, restrictToWorkspace, denyPatternSets...) hookPath := filepath.Join(tmpDir, fmt.Sprintf("hook-%s.sh", id)) if err := os.WriteFile(hookPath, []byte(hookScript), 0755); err != nil { return "", nil, fmt.Errorf("write hook script: %w", err) @@ -82,7 +83,7 @@ func generateSettingsJSON(hookPath string) []byte { } // generateHookScript creates a bash script that enforces GoClaw security policies. -func generateHookScript(workspace string, restrictToWorkspace bool) string { +func generateHookScript(workspace string, restrictToWorkspace bool, denyPatternSets ...[]*regexp.Regexp) string { var sb strings.Builder sb.WriteString(`#!/bin/bash @@ -115,7 +116,7 @@ check_shell_deny() { local patterns=( `) - for _, p := range ShellDenyPatterns { + for _, p := range hookDenyPatternStrings(denyPatternSets...) { // Escape single quotes for bash escaped := strings.ReplaceAll(p, `'`, `'\''`) fmt.Fprintf(&sb, " '%s'\n", escaped) @@ -205,3 +206,17 @@ allow return sb.String() } + +func hookDenyPatternStrings(denyPatternSets ...[]*regexp.Regexp) []string { + if len(denyPatternSets) == 0 { + return ShellDenyPatterns + } + patterns := make([]string, 0, len(denyPatternSets[0])) + for _, p := range denyPatternSets[0] { + if p == nil { + continue + } + patterns = append(patterns, p.String()) + } + return patterns +} diff --git a/internal/providers/claude_cli_hooks_test.go b/internal/providers/claude_cli_hooks_test.go new file mode 100644 index 00000000..3aa85f8a --- /dev/null +++ b/internal/providers/claude_cli_hooks_test.go @@ -0,0 +1,39 @@ +package providers + +import ( + "regexp" + "strings" + "testing" +) + +func TestGenerateHookScriptUsesConfiguredDenyPatterns(t *testing.T) { + pattern := regexp.MustCompile(`\bpip3?\s+install\b`) + + script := generateHookScript("", true, []*regexp.Regexp{pattern}) + + if !strings.Contains(script, pattern.String()) { + t.Fatalf("hook script missing configured pattern %q", pattern.String()) + } + if strings.Contains(script, `^\s*env\s*$`) { + t.Fatalf("hook script included default env_dump pattern when configured patterns were supplied") + } +} + +func TestGenerateHookScriptAllowsConfiguredEmptyDenyPatterns(t *testing.T) { + script := generateHookScript("", true, []*regexp.Regexp{}) + + if strings.Contains(script, `^\s*env\s*$`) { + t.Fatalf("hook script included default env_dump pattern for explicitly empty configured patterns") + } + if strings.Contains(script, `\bpip3?\s+install\b`) { + t.Fatalf("hook script included package-install pattern for explicitly empty configured patterns") + } +} + +func TestGenerateHookScriptDefaultsWhenNoConfiguredDenyPatterns(t *testing.T) { + script := generateHookScript("", true) + + if !strings.Contains(script, `^\s*env\s*$`) { + t.Fatalf("hook script missing default env_dump pattern") + } +}