diff --git a/README.md b/README.md index 859009b..ed631e9 100644 --- a/README.md +++ b/README.md @@ -235,6 +235,10 @@ Run the Vault MCP server: ```bash docker run --network=mcp -p 8080:8080 -e VAULT_ADDR='http://vault-dev:8200' -e VAULT_TOKEN='' -e TRANSPORT_MODE='http' vault-mcp-server:dev + +# Filter tools (optional) +docker run -i --rm vault-mcp-server:dev --toolsets=sys,kv +docker run -i --rm vault-mcp-server:dev --tools=read_secret,list_mounts ``` ## Available Tools @@ -332,6 +336,20 @@ Issues a new certificate using a PKI role. - `ipSans`: (Optional) IP SANs for the certificate - `ttl`: (Optional) Time-to-live for the certificate +### Tool Filtering + +Control which tools are available using `--toolsets` (groups) or `--tools` (individual): + +```bash +# Enable tool groups (default: all) +./vault-mcp-server --toolsets=sys,kv + +# Enable specific tools only +./vault-mcp-server --tools=read_secret,list_mounts,enable_pki +``` + +Available toolsets: `sys`, `kv`, `pki`, `all`, `default`. See `pkg/toolsets/mapping.go` for individual tool names. Cannot use both flags together. + ## Command Line Usage ```bash @@ -340,10 +358,10 @@ Issues a new certificate using a PKI role. # Run in stdio mode (default) ./vault-mcp-server -./vault-mcp-server stdio +./vault-mcp-server stdio [--log-file /path/to/log] [--toolsets ] [--tools ] # Run in HTTP mode -./vault-mcp-server http --transport-port 8080 --transport-host 127.0.0.1 +./vault-mcp-server http --transport-port 8080 --transport-host 127.0.0.1 [--toolsets ] [--tools ] # Show version ./vault-mcp-server --version @@ -414,6 +432,9 @@ vault-mcp-server/ │ │ ├── pki/ # PKI certificate tools │ │ ├── sys/ # System management tools │ │ └── tools.go # Tool registration +│ ├── toolsets/ # Toolset definitions and filtering +│ │ ├── toolsets.go # Toolset groups and helpers +│ │ └── mapping.go # Tool-to-toolset mapping │ └── utils/ # Utility functions ├── scripts/ # Build and utility scripts ├── version/ # Version information diff --git a/cmd/vault-mcp-server/init.go b/cmd/vault-mcp-server/init.go index c00637b..f782f2b 100644 --- a/cmd/vault-mcp-server/init.go +++ b/cmd/vault-mcp-server/init.go @@ -10,6 +10,7 @@ import ( stdlog "log" "os" + "github.com/hashicorp/vault-mcp-server/pkg/toolsets" "github.com/mark3labs/mcp-go/server" log "github.com/sirupsen/logrus" "github.com/spf13/cobra" @@ -31,6 +32,9 @@ func init() { httpCmdAlias.Flags().StringP("transport-port", "p", DefaultBindPort, "Port to listen on") httpCmdAlias.Flags().String("mcp-endpoint", DefaultEndPointPath, "Path for streamable HTTP endpoint") + rootCmd.PersistentFlags().String("toolsets", "all", toolsets.GenerateToolsetsHelp()) + rootCmd.PersistentFlags().String("tools", "", toolsets.GenerateToolsHelp()) + rootCmd.AddCommand(stdioCmd) rootCmd.AddCommand(streamableHTTPCmd) rootCmd.AddCommand(httpCmdAlias) // Add the alias for backward compatibility diff --git a/cmd/vault-mcp-server/main.go b/cmd/vault-mcp-server/main.go index 59e0189..1762e15 100644 --- a/cmd/vault-mcp-server/main.go +++ b/cmd/vault-mcp-server/main.go @@ -18,6 +18,7 @@ import ( "github.com/hashicorp/vault-mcp-server/pkg/client" "github.com/hashicorp/vault-mcp-server/pkg/tools" + "github.com/hashicorp/vault-mcp-server/pkg/toolsets" "github.com/hashicorp/vault-mcp-server/version" @@ -45,7 +46,7 @@ var ( Use: "stdio", Short: "Start stdio server", Long: `Start a server that communicates via standard input/output streams using JSON-RPC messages.`, - Run: func(_ *cobra.Command, _ []string) { + Run: func(cmd *cobra.Command, _ []string) { logFile, err := rootCmd.PersistentFlags().GetString("log-file") if err != nil { stdlog.Fatal("Failed to get log file:", err) @@ -55,7 +56,9 @@ var ( stdlog.Fatal("Failed to initialize logger:", err) } - if err := runStdioServer(logger); err != nil { + enabledToolsets := getToolsetsFromCmd(cmd.Root(), logger) + + if err := runStdioServer(logger, enabledToolsets); err != nil { stdlog.Fatal("failed to run stdio server:", err) } }, @@ -91,7 +94,9 @@ You can specify the host, port, and endpoint path to customize where the server stdlog.Fatal("Failed to get endpoint path:", err) } - if err := runHTTPServer(logger, host, port, endpointPath); err != nil { + enabledToolsets := getToolsetsFromCmd(cmd.Root(), logger) + + if err := runHTTPServer(logger, host, port, endpointPath, enabledToolsets); err != nil { stdlog.Fatal("failed to run streamableHTTP server:", err) } }, @@ -110,12 +115,12 @@ You can specify the host, port, and endpoint path to customize where the server } ) -func runHTTPServer(logger *log.Logger, host string, port string, endpointPath string) error { +func runHTTPServer(logger *log.Logger, host string, port string, endpointPath string, enabledToolsets []string) error { ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() hcServer := NewServer(version.Version, logger) - tools.InitTools(hcServer, logger) + tools.RegisterTools(hcServer, logger, enabledToolsets) return httpServerInit(ctx, hcServer, logger, host, port, endpointPath) } @@ -227,12 +232,12 @@ func httpServerInit(ctx context.Context, hcServer *server.MCPServer, logger *log return nil } -func runStdioServer(logger *log.Logger) error { +func runStdioServer(logger *log.Logger, enabledToolsets []string) error { ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) defer stop() hcServer := NewServer(version.Version, logger) - tools.InitTools(hcServer, logger) + tools.RegisterTools(hcServer, logger, enabledToolsets) return serverInit(ctx, hcServer, logger) } @@ -283,14 +288,16 @@ func runDefaultCommand(cmd *cobra.Command, _ []string) { stdlog.Fatal("Failed to initialize logger:", err) } - if err := runStdioServer(logger); err != nil { + enabledToolsets := getToolsetsFromCmd(cmd, logger) + + if err := runStdioServer(logger, enabledToolsets); err != nil { stdlog.Fatal("failed to run stdio server:", err) } } func main() { // Check environment variables first - they override command line args - if shouldUseHTTPMode() { + if shouldUseStreamableHTTPMode() { port := getHTTPPort() host := getHTTPHost() endpointPath := getEndpointPath(nil) @@ -301,8 +308,10 @@ func main() { stdlog.Fatal("Failed to initialize logger:", err) } - if err := runHTTPServer(logger, host, port, endpointPath); err != nil { - stdlog.Fatal("failed to run HTTP server:", err) + enabledToolsets := getToolsetsFromCmd(rootCmd, logger) + + if err := runHTTPServer(logger, host, port, endpointPath, enabledToolsets); err != nil { + stdlog.Fatal("failed to run StreamableHTTP server:", err) } return } @@ -314,8 +323,8 @@ func main() { } } -// shouldUseHTTPMode checks if environment variables indicate HTTP mode -func shouldUseHTTPMode() bool { +// shouldUseStreamableHTTPMode checks if environment variables indicate HTTP mode +func shouldUseStreamableHTTPMode() bool { transportMode := os.Getenv("TRANSPORT_MODE") return transportMode == "http" || transportMode == "streamable-http" || os.Getenv("TRANSPORT_PORT") != "" || @@ -339,7 +348,74 @@ func getHTTPHost() string { return DefaultBindAddress } -// Add function to get endpoint path from environment or flag +// parseToolsets parses and validates the toolsets flag value +func parseToolsets(toolsetsFlag string, logger *log.Logger) []string { + rawToolsets := strings.Split(toolsetsFlag, ",") + + cleaned, invalid := toolsets.CleanToolsets(rawToolsets) + if len(invalid) > 0 { + logger.Warnf("Invalid toolsets ignored: %v", invalid) + } + + expanded := toolsets.ExpandDefaultToolset(cleaned) + + logger.Infof("Enabled toolsets: %v", expanded) + return expanded +} + +// parseIndividualTools parses and validates the tools flag value +func parseIndividualTools(toolsFlag string, logger *log.Logger) []string { + rawTools := strings.Split(toolsFlag, ",") + + validTools, invalidTools := toolsets.ParseIndividualTools(rawTools) + if len(invalidTools) > 0 { + logger.Warnf("Invalid tool names ignored: %v", invalidTools) + } + + if len(validTools) == 0 { + logger.Warn("No valid tools specified, falling back to default toolsets") + return parseToolsets("default", logger) + } + + // Use the public API to enable individual tools mode + result := toolsets.EnableIndividualTools(validTools) + logger.Infof("Enabled individual tools: %v", validTools) + return result +} + +func getToolsetsFromCmd(cmd *cobra.Command, logger *log.Logger) []string { + // Check if --tools flag is set (individual tool mode) + toolsFlag, err := cmd.Flags().GetString("tools") + if err != nil { + // Try root persistent flags + toolsFlag, err = cmd.Root().PersistentFlags().GetString("tools") + } + + if err == nil && toolsFlag != "" { + // Ensure --toolsets is not also set + toolsetsFlag, _ := cmd.Flags().GetString("toolsets") + if toolsetsFlag == "" { + toolsetsFlag, _ = cmd.Root().PersistentFlags().GetString("toolsets") + } + if toolsetsFlag != "" && toolsetsFlag != "default" { + logger.Fatal("Cannot use both --tools and --toolsets flags together") + } + return parseIndividualTools(toolsFlag, logger) + } + + // Fall back to toolsets mode + toolsetsFlag, err := cmd.Flags().GetString("toolsets") + if err != nil { + toolsetsFlag, err = cmd.Root().PersistentFlags().GetString("toolsets") + if err != nil { + logger.Warnf("Failed to get toolsets flag, using default: %v", err) + toolsetsFlag = "default" + } + } + return parseToolsets(toolsetsFlag, logger) +} + +// getEndpointPath returns the endpoint path from environment or flag. func getEndpointPath(cmd *cobra.Command) string { // First check environment variable if envPath := os.Getenv("MCP_ENDPOINT"); envPath != "" { diff --git a/pkg/tools/tools.go b/pkg/tools/tools.go index 0a92838..05aa14a 100644 --- a/pkg/tools/tools.go +++ b/pkg/tools/tools.go @@ -7,60 +7,95 @@ import ( "github.com/hashicorp/vault-mcp-server/pkg/tools/kv" "github.com/hashicorp/vault-mcp-server/pkg/tools/pki" "github.com/hashicorp/vault-mcp-server/pkg/tools/sys" + "github.com/hashicorp/vault-mcp-server/pkg/toolsets" "github.com/mark3labs/mcp-go/server" log "github.com/sirupsen/logrus" ) -func InitTools(hcServer *server.MCPServer, logger *log.Logger) { +// RegisterTools registers all enabled tools on the MCP server. +// The enabledToolsets parameter controls which tools are registered. +func RegisterTools(hcServer *server.MCPServer, logger *log.Logger, enabledToolsets []string) { // Tools for Vault mount management - listMountsTool := sys.ListMounts(logger) - hcServer.AddTool(listMountsTool.Tool, listMountsTool.Handler) + if toolsets.IsToolEnabled("list_mounts", enabledToolsets) { + tool := sys.ListMounts(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } - createMountTool := sys.CreateMount(logger) - hcServer.AddTool(createMountTool.Tool, createMountTool.Handler) + if toolsets.IsToolEnabled("create_mount", enabledToolsets) { + tool := sys.CreateMount(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } - deleteMountTool := sys.DeleteMount(logger) - hcServer.AddTool(deleteMountTool.Tool, deleteMountTool.Handler) + if toolsets.IsToolEnabled("delete_mount", enabledToolsets) { + tool := sys.DeleteMount(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } // Tools for KV secrets management - listSecretsTool := kv.ListSecrets(logger) - hcServer.AddTool(listSecretsTool.Tool, listSecretsTool.Handler) - - readSecretTool := kv.ReadSecret(logger) - hcServer.AddTool(readSecretTool.Tool, readSecretTool.Handler) - - writeSecretTool := kv.WriteSecret(logger) - hcServer.AddTool(writeSecretTool.Tool, writeSecretTool.Handler) - - deleteSecretTool := kv.DeleteSecret(logger) - hcServer.AddTool(deleteSecretTool.Tool, deleteSecretTool.Handler) + if toolsets.IsToolEnabled("list_secrets", enabledToolsets) { + tool := kv.ListSecrets(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("read_secret", enabledToolsets) { + tool := kv.ReadSecret(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("write_secret", enabledToolsets) { + tool := kv.WriteSecret(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("delete_secret", enabledToolsets) { + tool := kv.DeleteSecret(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } // Tools for PKI management - enablePkiTool := pki.EnablePki(logger) - hcServer.AddTool(enablePkiTool.Tool, enablePkiTool.Handler) - - createPkiIssuer := pki.CreatePkiIssuer(logger) - hcServer.AddTool(createPkiIssuer.Tool, createPkiIssuer.Handler) - - listPkiIssuers := pki.ListPkiIssuers(logger) - hcServer.AddTool(listPkiIssuers.Tool, listPkiIssuers.Handler) - - readPkiIssuer := pki.ReadPkiIssuer(logger) - hcServer.AddTool(readPkiIssuer.Tool, readPkiIssuer.Handler) - - listPkiRoles := pki.ListPkiRoles(logger) - hcServer.AddTool(listPkiRoles.Tool, listPkiRoles.Handler) - - readPkiRole := pki.ReadPkiRole(logger) - hcServer.AddTool(readPkiRole.Tool, readPkiRole.Handler) - - createPkiRole := pki.CreatePkiRole(logger) - hcServer.AddTool(createPkiRole.Tool, createPkiRole.Handler) - - deletePkiRole := pki.DeletePkiRole(logger) - hcServer.AddTool(deletePkiRole.Tool, deletePkiRole.Handler) - - issuePkiCertificate := pki.IssuePkiCertificate(logger) - hcServer.AddTool(issuePkiCertificate.Tool, issuePkiCertificate.Handler) + if toolsets.IsToolEnabled("enable_pki", enabledToolsets) { + tool := pki.EnablePki(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("create_pki_issuer", enabledToolsets) { + tool := pki.CreatePkiIssuer(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("list_pki_issuers", enabledToolsets) { + tool := pki.ListPkiIssuers(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("read_pki_issuer", enabledToolsets) { + tool := pki.ReadPkiIssuer(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("list_pki_roles", enabledToolsets) { + tool := pki.ListPkiRoles(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("read_pki_role", enabledToolsets) { + tool := pki.ReadPkiRole(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("create_pki_role", enabledToolsets) { + tool := pki.CreatePkiRole(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("delete_pki_role", enabledToolsets) { + tool := pki.DeletePkiRole(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } + + if toolsets.IsToolEnabled("issue_pki_certificate", enabledToolsets) { + tool := pki.IssuePkiCertificate(logger) + hcServer.AddTool(tool.Tool, tool.Handler) + } } diff --git a/pkg/toolsets/mapping.go b/pkg/toolsets/mapping.go new file mode 100644 index 0000000..9710aee --- /dev/null +++ b/pkg/toolsets/mapping.go @@ -0,0 +1,97 @@ +// Copyright IBM Corp. 2025 +// SPDX-License-Identifier: MPL-2.0 + +package toolsets + +import ( + "slices" + "strings" +) + +// individualToolsMarker is an internal marker indicating individual tool mode. +const individualToolsMarker = "__individual_tools__" + +// ToolToToolset maps each tool name to its toolset. +var ToolToToolset = map[string]string{ + // Sys tools + "list_mounts": Sys, + "create_mount": Sys, + "delete_mount": Sys, + + // KV tools + "list_secrets": KV, + "read_secret": KV, + "write_secret": KV, + "delete_secret": KV, + "read_secret_metadata": KV, + "write_secret_metadata": KV, + "undelete_secret": KV, + "destroy_secret_versions": KV, + "patch_secret": KV, + + // PKI tools + "enable_pki": PKI, + "create_pki_issuer": PKI, + "list_pki_issuers": PKI, + "read_pki_issuer": PKI, + "list_pki_roles": PKI, + "read_pki_role": PKI, + "create_pki_role": PKI, + "delete_pki_role": PKI, + "issue_pki_certificate": PKI, +} + +// IsToolEnabled checks whether a tool should be registered given the enabled toolsets. +// If "all" is in the list, all tools are enabled. +// If the individualToolsMarker is present, only exact tool name matches are enabled. +// Otherwise, the tool is enabled if its parent toolset is in the enabled list. +func IsToolEnabled(toolName string, enabledToolsets []string) bool { + if slices.Contains(enabledToolsets, All) { + return true + } + + if slices.Contains(enabledToolsets, individualToolsMarker) { + return slices.Contains(enabledToolsets, toolName) + } + + toolset, ok := ToolToToolset[toolName] + if !ok { + return false + } + return slices.Contains(enabledToolsets, toolset) +} + +// GetAllValidToolNames returns a set of all known tool names. +func GetAllValidToolNames() map[string]bool { + names := make(map[string]bool, len(ToolToToolset)) + for name := range ToolToToolset { + names[name] = true + } + return names +} + +// ParseIndividualTools validates tool names and returns valid and invalid lists. +func ParseIndividualTools(tools []string) (valid, invalid []string) { + allNames := GetAllValidToolNames() + for _, name := range tools { + name = strings.TrimSpace(name) + if name == "" { + continue + } + if allNames[name] { + valid = append(valid, name) + } else { + invalid = append(invalid, name) + } + } + return valid, invalid +} + +// EnableIndividualTools creates an enabledToolsets slice for individual tool mode. +// It prepends the internal marker so IsToolEnabled knows to check exact names. +func EnableIndividualTools(tools []string) []string { + result := make([]string, 0, len(tools)+1) + result = append(result, individualToolsMarker) + result = append(result, tools...) + return result +} diff --git a/pkg/toolsets/mapping_test.go b/pkg/toolsets/mapping_test.go new file mode 100644 index 0000000..d8b073b --- /dev/null +++ b/pkg/toolsets/mapping_test.go @@ -0,0 +1,116 @@ +// Copyright IBM Corp. 2025 +// SPDX-License-Identifier: MPL-2.0 + +package toolsets + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestIsToolEnabled_All(t *testing.T) { + enabled := []string{"all"} + assert.True(t, IsToolEnabled("read_secret", enabled)) + assert.True(t, IsToolEnabled("list_mounts", enabled)) + assert.True(t, IsToolEnabled("enable_pki", enabled)) +} + +func TestIsToolEnabled_SpecificToolset(t *testing.T) { + enabled := []string{"kv"} + assert.True(t, IsToolEnabled("read_secret", enabled)) + assert.True(t, IsToolEnabled("list_secrets", enabled)) + assert.False(t, IsToolEnabled("enable_pki", enabled)) + assert.False(t, IsToolEnabled("list_mounts", enabled)) +} + +func TestIsToolEnabled_MultipleToolsets(t *testing.T) { + enabled := []string{"kv", "sys"} + assert.True(t, IsToolEnabled("read_secret", enabled)) + assert.True(t, IsToolEnabled("list_mounts", enabled)) + assert.False(t, IsToolEnabled("enable_pki", enabled)) +} + +func TestIsToolEnabled_IndividualTools(t *testing.T) { + enabled := EnableIndividualTools([]string{"read_secret", "list_secrets"}) + assert.True(t, IsToolEnabled("read_secret", enabled)) + assert.True(t, IsToolEnabled("list_secrets", enabled)) + assert.False(t, IsToolEnabled("write_secret", enabled)) + assert.False(t, IsToolEnabled("list_mounts", enabled)) +} + +func TestIsToolEnabled_Default(t *testing.T) { + enabled := ExpandDefaultToolset([]string{"default"}) + assert.True(t, IsToolEnabled("read_secret", enabled)) + assert.True(t, IsToolEnabled("list_mounts", enabled)) + assert.True(t, IsToolEnabled("enable_pki", enabled)) +} + +func TestIsToolEnabled_UnknownTool(t *testing.T) { + enabled := []string{"kv"} + assert.False(t, IsToolEnabled("nonexistent_tool", enabled)) +} + +func TestToolToToolset_Complete(t *testing.T) { + assert.Len(t, ToolToToolset, 21, "expected 21 tools mapped") + + // Verify sys tools + assert.Equal(t, Sys, ToolToToolset["list_mounts"]) + assert.Equal(t, Sys, ToolToToolset["create_mount"]) + assert.Equal(t, Sys, ToolToToolset["delete_mount"]) + + // Verify kv tools + kvTools := []string{ + "list_secrets", "read_secret", "write_secret", "delete_secret", + "read_secret_metadata", "write_secret_metadata", + "undelete_secret", "destroy_secret_versions", "patch_secret", + } + for _, tool := range kvTools { + assert.Equal(t, KV, ToolToToolset[tool], "tool %s should be in KV toolset", tool) + } + + // Verify pki tools + pkiTools := []string{ + "enable_pki", "create_pki_issuer", "list_pki_issuers", "read_pki_issuer", + "list_pki_roles", "read_pki_role", "create_pki_role", "delete_pki_role", + "issue_pki_certificate", + } + for _, tool := range pkiTools { + assert.Equal(t, PKI, ToolToToolset[tool], "tool %s should be in PKI toolset", tool) + } +} + +func TestParseIndividualTools(t *testing.T) { + t.Run("valid tools", func(t *testing.T) { + valid, invalid := ParseIndividualTools([]string{"read_secret", "list_mounts"}) + assert.Equal(t, []string{"read_secret", "list_mounts"}, valid) + assert.Empty(t, invalid) + }) + + t.Run("mix of valid and invalid", func(t *testing.T) { + valid, invalid := ParseIndividualTools([]string{"read_secret", "fake_tool"}) + assert.Equal(t, []string{"read_secret"}, valid) + assert.Equal(t, []string{"fake_tool"}, invalid) + }) + + t.Run("trims whitespace", func(t *testing.T) { + valid, invalid := ParseIndividualTools([]string{" read_secret ", ""}) + assert.Equal(t, []string{"read_secret"}, valid) + assert.Empty(t, invalid) + }) +} + +func TestGetAllValidToolNames(t *testing.T) { + names := GetAllValidToolNames() + require.Len(t, names, 21) + assert.True(t, names["read_secret"]) + assert.True(t, names["list_mounts"]) + assert.True(t, names["enable_pki"]) + assert.False(t, names["nonexistent"]) +} + +func TestEnableIndividualTools(t *testing.T) { + result := EnableIndividualTools([]string{"read_secret", "list_mounts"}) + assert.Equal(t, []string{individualToolsMarker, "read_secret", "list_mounts"}, result) +} diff --git a/pkg/toolsets/toolsets.go b/pkg/toolsets/toolsets.go new file mode 100644 index 0000000..57110ea --- /dev/null +++ b/pkg/toolsets/toolsets.go @@ -0,0 +1,131 @@ +// Copyright IBM Corp. 2025 +// SPDX-License-Identifier: MPL-2.0 + +package toolsets + +import ( + "fmt" + "slices" + "strings" +) + +const ( + // Sys represents mount management tools (list/create/delete mounts). + Sys = "sys" + // KV represents KV secrets tools (read/write/delete/patch/metadata/etc.). + KV = "kv" + // PKI represents PKI management tools (issuers/roles/certificates). + PKI = "pki" + + // All activates all toolsets. + All = "all" + // Default activates the default set of toolsets. + Default = "default" +) + +// Toolset describes a named group of tools. +type Toolset struct { + Name string + Description string +} + +// availableToolsets lists the concrete toolsets. +var availableToolsets = []Toolset{ + {Name: Sys, Description: "Mount management (list/create/delete mounts)"}, + {Name: KV, Description: "KV secrets (read/write/delete/patch/metadata/etc.)"}, + {Name: PKI, Description: "PKI management (issuers/roles/certificates)"}, +} + +// AvailableToolsets returns the list of concrete toolsets. +func AvailableToolsets() []Toolset { + return availableToolsets +} + +// DefaultToolsets returns the default set of enabled toolset names. +func DefaultToolsets() []string { + return []string{Sys, KV, PKI} +} + +// GetValidToolsetNames returns all valid toolset names including special keywords. +func GetValidToolsetNames() map[string]bool { + valid := map[string]bool{ + All: true, + Default: true, + } + for _, ts := range availableToolsets { + valid[ts.Name] = true + } + return valid +} + +// CleanToolsets deduplicates, trims, and validates toolset names. +// It returns the cleaned list and any invalid names found. +func CleanToolsets(input []string) (cleaned, invalid []string) { + seen := make(map[string]bool) + valid := GetValidToolsetNames() + + for _, name := range input { + name = strings.TrimSpace(name) + if name == "" { + continue + } + if !valid[name] { + invalid = append(invalid, name) + continue + } + if !seen[name] { + seen[name] = true + cleaned = append(cleaned, name) + } + } + return cleaned, invalid +} + +// ExpandDefaultToolset replaces the "default" keyword with the actual default toolsets. +func ExpandDefaultToolset(input []string) []string { + var result []string + seen := make(map[string]bool) + + for _, name := range input { + if name == Default { + for _, d := range DefaultToolsets() { + if !seen[d] { + seen[d] = true + result = append(result, d) + } + } + } else { + if !seen[name] { + seen[name] = true + result = append(result, name) + } + } + } + return result +} + +// ContainsToolset checks if a toolset name is present in the list. +func ContainsToolset(toolsets []string, name string) bool { + return slices.Contains(toolsets, name) +} + +// GenerateToolsetsHelp returns help text for the --toolsets flag. +func GenerateToolsetsHelp() string { + var sb strings.Builder + sb.WriteString("Comma-separated list of toolsets to enable.\n") + sb.WriteString("Special values: 'all' (enable all), 'default' (enable default set).\n") + sb.WriteString("Available toolsets:\n") + for _, ts := range availableToolsets { + fmt.Fprintf(&sb, " - %s: %s\n", ts.Name, ts.Description) + } + return sb.String() +} + +// GenerateToolsHelp returns help text for the --tools flag. +func GenerateToolsHelp() string { + var sb strings.Builder + sb.WriteString("Comma-separated list of individual tool names to enable.\n") + sb.WriteString("When specified, only the listed tools are enabled (overrides --toolsets).\n") + sb.WriteString("Use the tool names as they appear in the MCP tool list.\n") + return sb.String() +} diff --git a/pkg/toolsets/toolsets_test.go b/pkg/toolsets/toolsets_test.go new file mode 100644 index 0000000..da54200 --- /dev/null +++ b/pkg/toolsets/toolsets_test.go @@ -0,0 +1,103 @@ +// Copyright IBM Corp. 2025 +// SPDX-License-Identifier: MPL-2.0 + +package toolsets + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestDefaultToolsets(t *testing.T) { + defaults := DefaultToolsets() + assert.Contains(t, defaults, Sys) + assert.Contains(t, defaults, KV) + assert.Contains(t, defaults, PKI) + assert.Len(t, defaults, 3) +} + +func TestCleanToolsets(t *testing.T) { + t.Run("deduplicates and trims", func(t *testing.T) { + cleaned, invalid := CleanToolsets([]string{" kv ", "kv", "sys", " sys"}) + assert.Equal(t, []string{"kv", "sys"}, cleaned) + assert.Empty(t, invalid) + }) + + t.Run("detects invalid names", func(t *testing.T) { + cleaned, invalid := CleanToolsets([]string{"kv", "bogus", "nope"}) + assert.Equal(t, []string{"kv"}, cleaned) + assert.Equal(t, []string{"bogus", "nope"}, invalid) + }) + + t.Run("skips empty strings", func(t *testing.T) { + cleaned, invalid := CleanToolsets([]string{"", " ", "sys"}) + assert.Equal(t, []string{"sys"}, cleaned) + assert.Empty(t, invalid) + }) + + t.Run("accepts special keywords", func(t *testing.T) { + cleaned, invalid := CleanToolsets([]string{"all", "default"}) + assert.Equal(t, []string{"all", "default"}, cleaned) + assert.Empty(t, invalid) + }) +} + +func TestExpandDefaultToolset(t *testing.T) { + t.Run("expands default", func(t *testing.T) { + result := ExpandDefaultToolset([]string{"default"}) + assert.Equal(t, []string{"sys", "kv", "pki"}, result) + }) + + t.Run("preserves non-default", func(t *testing.T) { + result := ExpandDefaultToolset([]string{"kv"}) + assert.Equal(t, []string{"kv"}, result) + }) + + t.Run("deduplicates after expansion", func(t *testing.T) { + result := ExpandDefaultToolset([]string{"kv", "default"}) + assert.Equal(t, []string{"kv", "sys", "pki"}, result) + }) +} + +func TestContainsToolset(t *testing.T) { + assert.True(t, ContainsToolset([]string{"sys", "kv"}, "kv")) + assert.False(t, ContainsToolset([]string{"sys", "kv"}, "pki")) + assert.False(t, ContainsToolset(nil, "kv")) +} + +func TestGetValidToolsetNames(t *testing.T) { + valid := GetValidToolsetNames() + require.Len(t, valid, 5) + assert.True(t, valid[Sys]) + assert.True(t, valid[KV]) + assert.True(t, valid[PKI]) + assert.True(t, valid[All]) + assert.True(t, valid[Default]) +} + +func TestAvailableToolsets(t *testing.T) { + ts := AvailableToolsets() + assert.Len(t, ts, 3) + names := make([]string, len(ts)) + for i, toolset := range ts { + names[i] = toolset.Name + } + assert.Contains(t, names, Sys) + assert.Contains(t, names, KV) + assert.Contains(t, names, PKI) +} + +func TestGenerateToolsetsHelp(t *testing.T) { + help := GenerateToolsetsHelp() + assert.Contains(t, help, "sys") + assert.Contains(t, help, "kv") + assert.Contains(t, help, "pki") + assert.Contains(t, help, "all") +} + +func TestGenerateToolsHelp(t *testing.T) { + help := GenerateToolsHelp() + assert.Contains(t, help, "individual tool names") +}