From fb2c6e331bcaf98e740a4fdd70f4b52ff8b84eea Mon Sep 17 00:00:00 2001 From: Guoxun Wei Date: Wed, 17 Dec 2025 15:05:40 +0800 Subject: [PATCH 1/4] draft run cmd --- internal/azureclient/token_credential.go | 27 ++++ .../components/fleet/kubernetes/client.go | 3 +- internal/config/config.go | 37 ++++-- internal/k8s/adapter.go | 24 +++- internal/k8s/adapter_test.go | 6 +- internal/k8s/runcommand_executor.go | 121 ++++++++++++++++++ internal/server/server.go | 39 ++++-- 7 files changed, 228 insertions(+), 29 deletions(-) create mode 100644 internal/azureclient/token_credential.go create mode 100644 internal/k8s/runcommand_executor.go diff --git a/internal/azureclient/token_credential.go b/internal/azureclient/token_credential.go new file mode 100644 index 00000000..c145c6c8 --- /dev/null +++ b/internal/azureclient/token_credential.go @@ -0,0 +1,27 @@ +package azureclient + +import ( + "context" + "time" + + "github.com/Azure/azure-sdk-for-go/sdk/azcore" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" +) + +type StaticTokenCredential struct { + token string +} + +func NewStaticTokenCredential(token string) *StaticTokenCredential { + return &StaticTokenCredential{ + token: token, + } +} + +func (c *StaticTokenCredential) GetToken(ctx context.Context, opts policy.TokenRequestOptions) (azcore.AccessToken, error) { + expiresOn := time.Now().Add(1 * time.Hour) + return azcore.AccessToken{ + Token: c.token, + ExpiresOn: expiresOn, + }, nil +} diff --git a/internal/components/fleet/kubernetes/client.go b/internal/components/fleet/kubernetes/client.go index 4a44c165..6f6c6678 100644 --- a/internal/components/fleet/kubernetes/client.go +++ b/internal/components/fleet/kubernetes/client.go @@ -21,7 +21,8 @@ func NewClient() (*Client, error) { k8sExecutor := kubectl.NewExecutor() // Wrap it using the adapter to work with aks-mcp config - wrappedExecutor := k8s.WrapK8sExecutor(k8sExecutor) + // Fleet operations don't use multi-cluster mode (always use local kubeconfig) + wrappedExecutor := k8s.WrapK8sExecutor(k8sExecutor, false) return &Client{ executor: wrappedExecutor, diff --git a/internal/config/config.go b/internal/config/config.go index ae1e8e70..2a1c8819 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -74,22 +74,28 @@ type ConfigData struct { // Default is false (use new unified tools) // This flag is provided for backward compatibility and may be removed in future versions UseLegacyTools bool + + // EnableMultiCluster enables multi-cluster mode for kubectl tools + // When enabled, kubectl commands are executed via Azure AKS RunCommand API + // When disabled (default), kubectl commands are executed locally via kubeconfig + EnableMultiCluster bool } // NewConfig creates and returns a new configuration instance func NewConfig() *ConfigData { return &ConfigData{ - Timeout: 60, - CacheTimeout: 1 * time.Minute, - SecurityConfig: security.NewSecurityConfig(), - OAuthConfig: auth.NewDefaultOAuthConfig(), - Transport: "stdio", - Port: 8000, - AccessLevel: "readonly", - EnabledComponents: []string{}, - AllowNamespaces: "", - LogLevel: "info", - UseLegacyTools: os.Getenv("USE_LEGACY_TOOLS") == "true", + Timeout: 60, + CacheTimeout: 1 * time.Minute, + SecurityConfig: security.NewSecurityConfig(), + OAuthConfig: auth.NewDefaultOAuthConfig(), + Transport: "stdio", + Port: 8000, + AccessLevel: "readonly", + EnabledComponents: []string{}, + AllowNamespaces: "", + LogLevel: "info", + UseLegacyTools: os.Getenv("USE_LEGACY_TOOLS") == "true", + EnableMultiCluster: false, } } @@ -125,6 +131,10 @@ func (cfg *ConfigData) ParseFlags() { flag.StringVar(&cfg.AllowNamespaces, "allow-namespaces", "", "Comma-separated list of allowed Kubernetes namespaces (empty means all namespaces)") + // Multi-cluster configuration + flag.BoolVar(&cfg.EnableMultiCluster, "enable-multi-cluster", false, + "Enable multi-cluster mode for kubectl (uses Azure AKS RunCommand API instead of local kubeconfig)") + // Logging settings flag.StringVar(&cfg.LogLevel, "log-level", "info", "Log level (debug, info, warn, error)") @@ -268,6 +278,11 @@ func (cfg *ConfigData) ValidateConfig() error { return fmt.Errorf("OAuth authentication is not supported with stdio transport per MCP specification") } + // Validate multi-cluster + legacy tools compatibility + if cfg.EnableMultiCluster && cfg.UseLegacyTools { + return fmt.Errorf("multi-cluster mode (--enable-multi-cluster) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)") + } + return nil } diff --git a/internal/k8s/adapter.go b/internal/k8s/adapter.go index 2006f787..b90bf3d0 100644 --- a/internal/k8s/adapter.go +++ b/internal/k8s/adapter.go @@ -60,19 +60,35 @@ func ConvertConfig(cfg *config.ConfigData) *k8sconfig.ConfigData { // WrapK8sExecutor makes an mcp-kubernetes CommandExecutor // compatible with the aks-mcp tools.CommandExecutor interface. -func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor) tools.CommandExecutor { - return &executorAdapter{k8sExecutor: k8sExecutor} +func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor, enableMultiCluster bool) tools.CommandExecutor { + if enableMultiCluster { + return &executorAdapter{ + k8sExecutor: k8sExecutor, + runCommandExecutor: NewRunCommandExecutor(), + enableMultiCluster: true, + } + } + return &executorAdapter{ + k8sExecutor: k8sExecutor, + enableMultiCluster: false, + } } // executorAdapter bridges aks-mcp execution to mcp-kubernetes. // Unexported; behavior is defined by the wrapped executor. type executorAdapter struct { - k8sExecutor k8stools.CommandExecutor + k8sExecutor k8stools.CommandExecutor + runCommandExecutor *RunCommandExecutor + enableMultiCluster bool } // Execute adapts aks-mcp execution by converting its config -// and delegating to the wrapped mcp-kubernetes executor. +// and delegating to the wrapped mcp-kubernetes executor or RunCommand executor. func (a *executorAdapter) Execute(ctx context.Context, params map[string]interface{}, cfg *config.ConfigData) (string, error) { + if a.enableMultiCluster { + k8sCfg := ConvertConfig(cfg) + return a.runCommandExecutor.Execute(ctx, params, k8sCfg) + } k8sCfg := ConvertConfig(cfg) return a.k8sExecutor.Execute(ctx, params, k8sCfg) } diff --git a/internal/k8s/adapter_test.go b/internal/k8s/adapter_test.go index 9ec38d74..2509a357 100644 --- a/internal/k8s/adapter_test.go +++ b/internal/k8s/adapter_test.go @@ -154,7 +154,7 @@ func TestExecutorAdapter_DelegatesAndForwards(t *testing.T) { t.Parallel() fe := &fakeExecutor{out: "ok"} - adapter := WrapK8sExecutor(fe) + adapter := WrapK8sExecutor(fe, false) params := map[string]interface{}{"k": "v"} inCfg := &config.ConfigData{ @@ -191,7 +191,7 @@ func TestExecutorAdapter_PropagatesError(t *testing.T) { t.Parallel() fe := &fakeExecutor{err: errors.New("boom")} - adapter := WrapK8sExecutor(fe) + adapter := WrapK8sExecutor(fe, false) _, err := adapter.Execute(context.Background(), map[string]interface{}{"x": 1}, &config.ConfigData{}) if err == nil { @@ -210,7 +210,7 @@ func TestExecutorAdapter_PanicsOnNilConfig_CurrentBehavior(t *testing.T) { }() fe := &fakeExecutor{} - adapter := WrapK8sExecutor(fe) + adapter := WrapK8sExecutor(fe, false) _, _ = adapter.Execute(context.Background(), map[string]interface{}{"x": 1}, nil) } diff --git a/internal/k8s/runcommand_executor.go b/internal/k8s/runcommand_executor.go new file mode 100644 index 00000000..8002fcc6 --- /dev/null +++ b/internal/k8s/runcommand_executor.go @@ -0,0 +1,121 @@ +package k8s + +import ( + "context" + "encoding/json" + "fmt" + "strings" + + "github.com/Azure/aks-mcp/internal/azureclient" + "github.com/Azure/aks-mcp/internal/logger" + "github.com/Azure/azure-sdk-for-go/sdk/azcore/arm" + "github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/containerservice/armcontainerservice/v2" + k8sconfig "github.com/Azure/mcp-kubernetes/pkg/config" +) + +type RunCommandExecutor struct{} + +func NewRunCommandExecutor() *RunCommandExecutor { + return &RunCommandExecutor{} +} + +type RequestContext struct { + AzureToken string `json:"azure_token"` + SubscriptionID string `json:"subscription_id"` + ResourceGroup string `json:"resource_group"` + ClusterName string `json:"cluster_name"` +} + +func extractRequestContext(ctx context.Context) (*RequestContext, error) { + requestContextStr, ok := ctx.Value("request_context").(string) + if !ok || requestContextStr == "" { + return nil, fmt.Errorf("request_context not found in context or empty") + } + + var reqCtx RequestContext + if err := json.Unmarshal([]byte(requestContextStr), &reqCtx); err != nil { + return nil, fmt.Errorf("failed to unmarshal request_context: %w", err) + } + + if reqCtx.AzureToken == "" { + return nil, fmt.Errorf("azure_token is required in request_context") + } + if reqCtx.SubscriptionID == "" { + return nil, fmt.Errorf("subscription_id is required in request_context") + } + if reqCtx.ResourceGroup == "" { + return nil, fmt.Errorf("resource_group is required in request_context") + } + if reqCtx.ClusterName == "" { + return nil, fmt.Errorf("cluster_name is required in request_context") + } + + return &reqCtx, nil +} + +func (e *RunCommandExecutor) Execute(ctx context.Context, params map[string]interface{}, cfg *k8sconfig.ConfigData) (string, error) { + logger.Debugf("RunCommandExecutor: Executing kubectl command via AKS RunCommand API") + + reqCtx, err := extractRequestContext(ctx) + if err != nil { + return "", fmt.Errorf("failed to extract request context: %w", err) + } + + logger.Debugf("RunCommandExecutor: Using cluster %s in resource group %s", reqCtx.ClusterName, reqCtx.ResourceGroup) + + command, err := e.buildCommand(params) + if err != nil { + return "", fmt.Errorf("failed to build command: %w", err) + } + + logger.Debugf("RunCommandExecutor: Command to execute: %s", command) + + cred := azureclient.NewStaticTokenCredential(reqCtx.AzureToken) + + clientFactory, err := armcontainerservice.NewClientFactory(reqCtx.SubscriptionID, cred, &arm.ClientOptions{}) + if err != nil { + return "", fmt.Errorf("failed to create Azure client for cluster %s/%s, command '%s': %w", reqCtx.ResourceGroup, reqCtx.ClusterName, command, err) + } + + managedClustersClient := clientFactory.NewManagedClustersClient() + + runCommandRequest := armcontainerservice.RunCommandRequest{ + Command: &command, + } + + poller, err := managedClustersClient.BeginRunCommand(ctx, reqCtx.ResourceGroup, reqCtx.ClusterName, runCommandRequest, nil) + if err != nil { + return "", fmt.Errorf("failed to start run command '%s' on cluster %s/%s: %w", command, reqCtx.ResourceGroup, reqCtx.ClusterName, err) + } + + logger.Debugf("RunCommandExecutor: Waiting for command to complete...") + resp, err := poller.PollUntilDone(ctx, nil) + if err != nil { + return "", fmt.Errorf("failed to execute command '%s' on cluster %s/%s: %w", command, reqCtx.ResourceGroup, reqCtx.ClusterName, err) + } + + if resp.Properties == nil { + return "", fmt.Errorf("run command '%s' on cluster %s/%s returned nil properties", command, reqCtx.ResourceGroup, reqCtx.ClusterName) + } + + var output strings.Builder + if resp.Properties.Logs != nil { + output.WriteString(*resp.Properties.Logs) + } + + // Check if command execution failed + if resp.Properties.ExitCode != nil && *resp.Properties.ExitCode != 0 { + return output.String(), fmt.Errorf("command '%s' on cluster %s/%s failed with exit code %d", command, reqCtx.ResourceGroup, reqCtx.ClusterName, *resp.Properties.ExitCode) + } + + logger.Debugf("RunCommandExecutor: Command completed successfully") + return output.String(), nil +} + +func (e *RunCommandExecutor) buildCommand(params map[string]interface{}) (string, error) { + cmd, ok := params["command"].(string) + if !ok { + return "", fmt.Errorf("command parameter is required") + } + return cmd, nil +} diff --git a/internal/server/server.go b/internal/server/server.go index 75d0c193..e7561db5 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -1,6 +1,7 @@ package server import ( + "context" "encoding/json" "fmt" "net/http" @@ -32,7 +33,6 @@ import ( "github.com/Azure/mcp-kubernetes/pkg/helm" "github.com/Azure/mcp-kubernetes/pkg/hubble" "github.com/Azure/mcp-kubernetes/pkg/kubectl" - k8stools "github.com/Azure/mcp-kubernetes/pkg/tools" "github.com/mark3labs/mcp-go/server" ) @@ -369,8 +369,17 @@ func (s *Service) Run() error { case "sse": addr := fmt.Sprintf("%s:%d", s.cfg.Host, s.cfg.Port) - // Create SSE server first - sse := server.NewSSEServer(s.mcpServer) + // Create SSE server with context function to extract Azure token from headers + sse := server.NewSSEServer( + s.mcpServer, + server.WithSSEContextFunc(func(ctx context.Context, r *http.Request) context.Context { + // Extract request context from X-Request-Context header and add to context + if requestContext := r.Header.Get("X-Request-Context"); requestContext != "" { + ctx = context.WithValue(ctx, "request_context", requestContext) + } + return ctx + }), + ) // Create custom HTTP server with helpful 404 responses customServer := s.createCustomSSEServerWithHelp404(sse, addr) @@ -395,6 +404,13 @@ func (s *Service) Run() error { streamableServer := server.NewStreamableHTTPServer( s.mcpServer, server.WithStreamableHTTPServer(customServer), + server.WithHTTPContextFunc(func(ctx context.Context, r *http.Request) context.Context { + // Extract request context from X-Request-Context header and add to context + if requestContext := r.Header.Get("X-Request-Context"); requestContext != "" { + ctx = context.WithValue(ctx, "request_context", requestContext) + } + return ctx + }), ) // Update the mux to use the actual streamable server as the MCP handler @@ -509,14 +525,17 @@ func (s *Service) registerKubectlComponent() { // Create a kubectl executor kubectlExecutor := kubectl.NewKubectlToolExecutor() - // Convert aks-mcp config to k8s config - k8sCfg := k8s.ConvertConfig(s.cfg) + // Wrap the executor with multi-cluster support if enabled + wrappedExecutor := k8s.WrapK8sExecutor(kubectlExecutor, s.cfg.EnableMultiCluster) + if s.cfg.EnableMultiCluster { + logger.Infof("Multi-cluster mode enabled: kubectl commands will use Azure AKS RunCommand API") + } // Register each kubectl tool for _, tool := range kubectlTools { logger.Debugf("Registering kubectl tool: %s", tool.Name) - // Create a handler that injects the tool name into params - handler := k8stools.CreateToolHandlerWithName(kubectlExecutor, k8sCfg, tool.Name) + // Create a handler that uses our wrapped executor + handler := tools.CreateToolHandler(wrappedExecutor, s.cfg) s.mcpServer.AddTool(tool, handler) } } @@ -640,7 +659,7 @@ func (s *Service) registerDetectorComponent() { func (s *Service) registerHelmComponent() { logger.Debugf("Registering Kubernetes tool: helm") helmTool := helm.RegisterHelm() - helmExecutor := k8s.WrapK8sExecutor(helm.NewExecutor()) + helmExecutor := k8s.WrapK8sExecutor(helm.NewExecutor(), s.cfg.EnableMultiCluster) s.mcpServer.AddTool(helmTool, tools.CreateToolHandler(helmExecutor, s.cfg)) } @@ -648,7 +667,7 @@ func (s *Service) registerHelmComponent() { func (s *Service) registerCiliumComponent() { logger.Debugf("Registering Kubernetes tool: cilium") ciliumTool := cilium.RegisterCilium() - ciliumExecutor := k8s.WrapK8sExecutor(cilium.NewExecutor()) + ciliumExecutor := k8s.WrapK8sExecutor(cilium.NewExecutor(), s.cfg.EnableMultiCluster) s.mcpServer.AddTool(ciliumTool, tools.CreateToolHandler(ciliumExecutor, s.cfg)) } @@ -656,7 +675,7 @@ func (s *Service) registerCiliumComponent() { func (s *Service) registerHubbleComponent() { logger.Debugf("Registering Kubernetes tool: hubble") hubbleTool := hubble.RegisterHubble() - hubbleExecutor := k8s.WrapK8sExecutor(hubble.NewExecutor()) + hubbleExecutor := k8s.WrapK8sExecutor(hubble.NewExecutor(), s.cfg.EnableMultiCluster) s.mcpServer.AddTool(hubbleTool, tools.CreateToolHandler(hubbleExecutor, s.cfg)) } From 2ca77cfcc7e6edffc9a52fad789dab145142c73f Mon Sep 17 00:00:00 2001 From: Guoxun Wei Date: Wed, 21 Jan 2026 13:50:33 +0800 Subject: [PATCH 2/4] use header --- .../components/fleet/kubernetes/client.go | 2 +- internal/ctx/context.go | 5 + internal/k8s/adapter.go | 10 +- internal/k8s/registry.go | 94 +++++++++++++++++++ internal/k8s/runcommand_executor.go | 36 ++++--- internal/server/server.go | 51 +++++----- 6 files changed, 155 insertions(+), 43 deletions(-) create mode 100644 internal/ctx/context.go create mode 100644 internal/k8s/registry.go diff --git a/internal/components/fleet/kubernetes/client.go b/internal/components/fleet/kubernetes/client.go index 6f6c6678..1851f0f0 100644 --- a/internal/components/fleet/kubernetes/client.go +++ b/internal/components/fleet/kubernetes/client.go @@ -21,7 +21,7 @@ func NewClient() (*Client, error) { k8sExecutor := kubectl.NewExecutor() // Wrap it using the adapter to work with aks-mcp config - // Fleet operations don't use multi-cluster mode (always use local kubeconfig) + // Fleet operations don't use multi-cluster mode for now (always use local kubeconfig) wrappedExecutor := k8s.WrapK8sExecutor(k8sExecutor, false) return &Client{ diff --git a/internal/ctx/context.go b/internal/ctx/context.go new file mode 100644 index 00000000..c7204015 --- /dev/null +++ b/internal/ctx/context.go @@ -0,0 +1,5 @@ +package ctx + +type ContextKey string + +const AzureTokenKey ContextKey = "X-Azure-Token" diff --git a/internal/k8s/adapter.go b/internal/k8s/adapter.go index b90bf3d0..9dae8828 100644 --- a/internal/k8s/adapter.go +++ b/internal/k8s/adapter.go @@ -61,16 +61,10 @@ func ConvertConfig(cfg *config.ConfigData) *k8sconfig.ConfigData { // WrapK8sExecutor makes an mcp-kubernetes CommandExecutor // compatible with the aks-mcp tools.CommandExecutor interface. func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor, enableMultiCluster bool) tools.CommandExecutor { - if enableMultiCluster { - return &executorAdapter{ - k8sExecutor: k8sExecutor, - runCommandExecutor: NewRunCommandExecutor(), - enableMultiCluster: true, - } - } return &executorAdapter{ k8sExecutor: k8sExecutor, - enableMultiCluster: false, + runCommandExecutor: NewRunCommandExecutor(), + enableMultiCluster: enableMultiCluster, } } diff --git a/internal/k8s/registry.go b/internal/k8s/registry.go new file mode 100644 index 00000000..c1afcadc --- /dev/null +++ b/internal/k8s/registry.go @@ -0,0 +1,94 @@ +package k8s + +import ( + "fmt" + "strings" + + "github.com/Azure/mcp-kubernetes/pkg/kubectl" + "github.com/Azure/mcp-kubernetes/pkg/security" + "github.com/mark3labs/mcp-go/mcp" +) + +const ( + AccessLevelReadOnly = "readonly" + AccessLevelReadWrite = "readwrite" +) + +func createCallKubectlTool(accessLevel string) mcp.Tool { + var description string + + readCommands := strings.Join(security.KubectlReadOperations, ", ") + writeCommands := strings.Join(security.KubectlReadWriteOperations, ", ") + + switch accessLevel { + case AccessLevelReadOnly: + description = fmt.Sprintf(`Execute kubectl commands with read-only access. + +Pass full kubectl command including 'kubectl' prefix. All standard kubectl flags are supported. + +Allowed commands: +%s + +Examples: +- command='kubectl get pods -n default' +- command='kubectl describe deployment myapp -n production' +- command='kubectl logs nginx-pod -f' +- command='kubectl top pods' +- command='kubectl events --all-namespaces' +- command='kubectl explain pods.spec.containers' +- command='kubectl auth can-i create pods'`, readCommands) + case AccessLevelReadWrite: + description = fmt.Sprintf(`Execute kubectl commands with read and write access. + +Pass full kubectl command including 'kubectl' prefix. All standard kubectl flags are supported. + +Allowed commands: +Read: %s +Write: %s + +Examples: +- command='kubectl get pods -n default' +- command='kubectl create -f deployment.yaml' +- command='kubectl apply -f deployment.yaml' +- command='kubectl delete pod nginx-pod' +- command='kubectl scale deployment myapp --replicas=3' +- command='kubectl rollout status deployment/myapp' +- command='kubectl label pods foo unhealthy=true' +- command='kubectl exec nginx-pod -- date' +- command='kubectl config use-context my-cluster-context'`, readCommands, writeCommands) + default: + description = fmt.Sprintf(`Execute kubectl commands with unknown access level (defaulting to read-only). + +Pass full kubectl command including 'kubectl' prefix. All standard kubectl flags are supported. + +Allowed commands: +%s + +Examples: +- command='kubectl get pods -n default' +- command='kubectl describe deployment myapp -n production' +- command='kubectl logs nginx-pod -f'`, readCommands) + } + + return mcp.NewTool("call_kubectl", + mcp.WithDescription(description), + mcp.WithString("command", + mcp.Required(), + mcp.Description("Full kubectl command to execute (e.g., 'kubectl get pods -n default', 'kubectl describe deployment myapp', 'kubectl logs nginx-pod -f')"), + ), + mcp.WithString("aks_resource_id", + mcp.Required(), + mcp.Description("Full Azure Resource ID of the AKS cluster (e.g., /subscriptions/{subscriptionId}/resourceGroups/{resourceGroupName}/providers/Microsoft.ContainerService/managedClusters/{clusterName})"), + ), + ) +} + +func RegisterKubectlTools(accessLevel string, useUnifiedTool bool, enableMultiCluster bool) []mcp.Tool { + if enableMultiCluster { + return []mcp.Tool{ + createCallKubectlTool(accessLevel), + } + } + + return kubectl.RegisterKubectlTools(accessLevel, useUnifiedTool) +} diff --git a/internal/k8s/runcommand_executor.go b/internal/k8s/runcommand_executor.go index 8002fcc6..baee9437 100644 --- a/internal/k8s/runcommand_executor.go +++ b/internal/k8s/runcommand_executor.go @@ -2,11 +2,11 @@ package k8s import ( "context" - "encoding/json" "fmt" "strings" "github.com/Azure/aks-mcp/internal/azureclient" + "github.com/Azure/aks-mcp/internal/ctx" "github.com/Azure/aks-mcp/internal/logger" "github.com/Azure/azure-sdk-for-go/sdk/azcore/arm" "github.com/Azure/azure-sdk-for-go/sdk/resourcemanager/containerservice/armcontainerservice/v2" @@ -26,17 +26,31 @@ type RequestContext struct { ClusterName string `json:"cluster_name"` } -func extractRequestContext(ctx context.Context) (*RequestContext, error) { - requestContextStr, ok := ctx.Value("request_context").(string) - if !ok || requestContextStr == "" { - return nil, fmt.Errorf("request_context not found in context or empty") +func extractRequestContext(c context.Context, params map[string]interface{}) (*RequestContext, error) { + tokenStr, ok := c.Value(ctx.AzureTokenKey).(string) + if !ok || tokenStr == "" { + return nil, fmt.Errorf("X-Azure-Token not found in context or empty") } var reqCtx RequestContext - if err := json.Unmarshal([]byte(requestContextStr), &reqCtx); err != nil { - return nil, fmt.Errorf("failed to unmarshal request_context: %w", err) + reqCtx.AzureToken = tokenStr + + // Extract aks_resource_id from params + aksResourceID, ok := params["aks_resource_id"].(string) + if !ok || aksResourceID == "" { + return nil, fmt.Errorf("aks_resource_id not found in params or empty") + } + + // Parse Azure Resource ID: /subscriptions/{sub}/resourceGroups/{rg}/providers/Microsoft.ContainerService/managedClusters/{cluster} + resourceID, err := arm.ParseResourceID(aksResourceID) + if err != nil { + return nil, fmt.Errorf("failed to parse aks_resource_id: %w", err) } + reqCtx.SubscriptionID = resourceID.SubscriptionID + reqCtx.ResourceGroup = resourceID.ResourceGroupName + reqCtx.ClusterName = resourceID.Name + if reqCtx.AzureToken == "" { return nil, fmt.Errorf("azure_token is required in request_context") } @@ -54,21 +68,18 @@ func extractRequestContext(ctx context.Context) (*RequestContext, error) { } func (e *RunCommandExecutor) Execute(ctx context.Context, params map[string]interface{}, cfg *k8sconfig.ConfigData) (string, error) { - logger.Debugf("RunCommandExecutor: Executing kubectl command via AKS RunCommand API") - reqCtx, err := extractRequestContext(ctx) + reqCtx, err := extractRequestContext(ctx, params) if err != nil { return "", fmt.Errorf("failed to extract request context: %w", err) } - logger.Debugf("RunCommandExecutor: Using cluster %s in resource group %s", reqCtx.ClusterName, reqCtx.ResourceGroup) - command, err := e.buildCommand(params) if err != nil { return "", fmt.Errorf("failed to build command: %w", err) } - logger.Debugf("RunCommandExecutor: Command to execute: %s", command) + logger.Debugf("RunCommandExecutor: Command to execute: %s in cluster %s/%s", command, reqCtx.ResourceGroup, reqCtx.ClusterName) cred := azureclient.NewStaticTokenCredential(reqCtx.AzureToken) @@ -88,7 +99,6 @@ func (e *RunCommandExecutor) Execute(ctx context.Context, params map[string]inte return "", fmt.Errorf("failed to start run command '%s' on cluster %s/%s: %w", command, reqCtx.ResourceGroup, reqCtx.ClusterName, err) } - logger.Debugf("RunCommandExecutor: Waiting for command to complete...") resp, err := poller.PollUntilDone(ctx, nil) if err != nil { return "", fmt.Errorf("failed to execute command '%s' on cluster %s/%s: %w", command, reqCtx.ResourceGroup, reqCtx.ClusterName, err) diff --git a/internal/server/server.go b/internal/server/server.go index e7561db5..28d4ea84 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -23,6 +23,7 @@ import ( "github.com/Azure/aks-mcp/internal/components/monitor" "github.com/Azure/aks-mcp/internal/components/network" "github.com/Azure/aks-mcp/internal/config" + "github.com/Azure/aks-mcp/internal/ctx" "github.com/Azure/aks-mcp/internal/k8s" "github.com/Azure/aks-mcp/internal/logger" "github.com/Azure/aks-mcp/internal/prompts" @@ -171,14 +172,18 @@ func (s *Service) registerAllComponents() { logger.Infof("All components enabled by default") } - // Azure Components - s.registerAzureComponents() - // Kubernetes Components s.registerKubernetesComponents() - // Prompts - s.registerPrompts() + if s.cfg.EnableMultiCluster { + logger.Infof("Multi-cluster mode enabled - skipping Azure component registration because they are not yet supported in multi-cluster mode") + } else { + // Azure Components + s.registerAzureComponents() + + // Prompts + s.registerPrompts() + } } // registerPrompts registers all available prompts @@ -372,12 +377,11 @@ func (s *Service) Run() error { // Create SSE server with context function to extract Azure token from headers sse := server.NewSSEServer( s.mcpServer, - server.WithSSEContextFunc(func(ctx context.Context, r *http.Request) context.Context { - // Extract request context from X-Request-Context header and add to context - if requestContext := r.Header.Get("X-Request-Context"); requestContext != "" { - ctx = context.WithValue(ctx, "request_context", requestContext) + server.WithSSEContextFunc(func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) } - return ctx + return c }), ) @@ -404,12 +408,12 @@ func (s *Service) Run() error { streamableServer := server.NewStreamableHTTPServer( s.mcpServer, server.WithStreamableHTTPServer(customServer), - server.WithHTTPContextFunc(func(ctx context.Context, r *http.Request) context.Context { - // Extract request context from X-Request-Context header and add to context - if requestContext := r.Header.Get("X-Request-Context"); requestContext != "" { - ctx = context.WithValue(ctx, "request_context", requestContext) + server.WithHTTPContextFunc(func(c context.Context, r *http.Request) context.Context { + // Extract request context from X-Azure-Token header and add to context + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) } - return ctx + return c }), ) @@ -501,8 +505,13 @@ func (s *Service) registerKubernetesComponents() { // Core Kubernetes Component (kubectl) s.registerKubectlComponent() - // Optional Kubernetes Components (based on configuration) - s.registerOptionalKubernetesComponents() + // Do not register optional components in multi-cluster mode, they are not supported yet. + if s.cfg.EnableMultiCluster { + logger.Infof("Multi-cluster mode enabled - skipping optional Kubernetes component registration because they are not yet supported in multi-cluster mode") + } else { + // Optional Kubernetes Components (based on configuration) + s.registerOptionalKubernetesComponents() + } logger.Infof("Kubernetes Components registered successfully") } @@ -520,7 +529,7 @@ func (s *Service) registerKubectlComponent() { } // Get kubectl tools filtered by access level and tool type - kubectlTools := kubectl.RegisterKubectlTools(s.cfg.AccessLevel, useUnifiedTool) + kubectlTools := k8s.RegisterKubectlTools(s.cfg.AccessLevel, useUnifiedTool, s.cfg.EnableMultiCluster) // Create a kubectl executor kubectlExecutor := kubectl.NewKubectlToolExecutor() @@ -659,7 +668,7 @@ func (s *Service) registerDetectorComponent() { func (s *Service) registerHelmComponent() { logger.Debugf("Registering Kubernetes tool: helm") helmTool := helm.RegisterHelm() - helmExecutor := k8s.WrapK8sExecutor(helm.NewExecutor(), s.cfg.EnableMultiCluster) + helmExecutor := k8s.WrapK8sExecutor(helm.NewExecutor(), false) s.mcpServer.AddTool(helmTool, tools.CreateToolHandler(helmExecutor, s.cfg)) } @@ -667,7 +676,7 @@ func (s *Service) registerHelmComponent() { func (s *Service) registerCiliumComponent() { logger.Debugf("Registering Kubernetes tool: cilium") ciliumTool := cilium.RegisterCilium() - ciliumExecutor := k8s.WrapK8sExecutor(cilium.NewExecutor(), s.cfg.EnableMultiCluster) + ciliumExecutor := k8s.WrapK8sExecutor(cilium.NewExecutor(), false) s.mcpServer.AddTool(ciliumTool, tools.CreateToolHandler(ciliumExecutor, s.cfg)) } @@ -675,7 +684,7 @@ func (s *Service) registerCiliumComponent() { func (s *Service) registerHubbleComponent() { logger.Debugf("Registering Kubernetes tool: hubble") hubbleTool := hubble.RegisterHubble() - hubbleExecutor := k8s.WrapK8sExecutor(hubble.NewExecutor(), s.cfg.EnableMultiCluster) + hubbleExecutor := k8s.WrapK8sExecutor(hubble.NewExecutor(), false) s.mcpServer.AddTool(hubbleTool, tools.CreateToolHandler(hubbleExecutor, s.cfg)) } From 7146efa7d336c910d4c45d835f94ef2b3d7c7515 Mon Sep 17 00:00:00 2001 From: Guoxun Wei Date: Wed, 21 Jan 2026 15:56:05 +0800 Subject: [PATCH 3/4] add unit test --- internal/config/config.go | 5 + internal/config/config_test.go | 256 +++++++++++++++++++++ internal/ctx/context.go | 3 + internal/k8s/runcommand_executor_test.go | 281 +++++++++++++++++++++++ internal/server/context_test.go | 211 +++++++++++++++++ 5 files changed, 756 insertions(+) create mode 100644 internal/k8s/runcommand_executor_test.go create mode 100644 internal/server/context_test.go diff --git a/internal/config/config.go b/internal/config/config.go index 2a1c8819..a529de8b 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -278,6 +278,11 @@ func (cfg *ConfigData) ValidateConfig() error { return fmt.Errorf("OAuth authentication is not supported with stdio transport per MCP specification") } + // Validate multi-cluster + stdio transport compatibility + if cfg.EnableMultiCluster && cfg.Transport == "stdio" { + return fmt.Errorf("multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport, use sse or streamable-http instead") + } + // Validate multi-cluster + legacy tools compatibility if cfg.EnableMultiCluster && cfg.UseLegacyTools { return fmt.Errorf("multi-cluster mode (--enable-multi-cluster) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)") diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 39345584..834e6ced 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -336,3 +336,259 @@ func setEnv(t *testing.T, key, value string) { func unsetEnv(t *testing.T, key string) { t.Setenv(key, "") } + +func TestValidateConfig_OAuthWithStdio(t *testing.T) { + cfg := NewConfig() + cfg.OAuthConfig.Enabled = true + cfg.Transport = "stdio" + + err := cfg.ValidateConfig() + if err == nil { + t.Fatal("Expected error when OAuth is enabled with stdio transport, got nil") + } + + expectedMsg := "OAuth authentication is not supported with stdio transport per MCP specification" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestValidateConfig_OAuthWithSSE(t *testing.T) { + cfg := NewConfig() + cfg.OAuthConfig.Enabled = true + cfg.Transport = "sse" + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for OAuth with SSE transport, got: %v", err) + } +} + +func TestValidateConfig_OAuthWithStreamableHTTP(t *testing.T) { + cfg := NewConfig() + cfg.OAuthConfig.Enabled = true + cfg.Transport = "streamable-http" + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for OAuth with streamable-http transport, got: %v", err) + } +} + +func TestValidateConfig_MultiClusterWithLegacyTools(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = true + cfg.UseLegacyTools = true + cfg.Transport = "sse" + + err := cfg.ValidateConfig() + if err == nil { + t.Fatal("Expected error when multi-cluster is enabled with legacy tools, got nil") + } + + expectedMsg := "multi-cluster mode (--enable-multi-cluster) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestValidateConfig_MultiClusterWithStdio(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = true + cfg.Transport = "stdio" + + err := cfg.ValidateConfig() + if err == nil { + t.Fatal("Expected error when multi-cluster is enabled with stdio transport, got nil") + } + + expectedMsg := "multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport, use sse or streamable-http instead" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestValidateConfig_MultiClusterWithSSE(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = true + cfg.Transport = "sse" + cfg.UseLegacyTools = false + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for multi-cluster with SSE transport, got: %v", err) + } +} + +func TestValidateConfig_MultiClusterWithStreamableHTTP(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = true + cfg.Transport = "streamable-http" + cfg.UseLegacyTools = false + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for multi-cluster with streamable-http transport, got: %v", err) + } +} + +func TestValidateConfig_MultiClusterWithUnifiedTools(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = true + cfg.UseLegacyTools = false + cfg.Transport = "sse" + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for multi-cluster with unified tools, got: %v", err) + } +} + +func TestValidateConfig_LegacyToolsWithoutMultiCluster(t *testing.T) { + cfg := NewConfig() + cfg.EnableMultiCluster = false + cfg.UseLegacyTools = true + + err := cfg.ValidateConfig() + if err != nil { + t.Errorf("Expected no error for legacy tools without multi-cluster, got: %v", err) + } +} + +func TestValidateConfig_ValidCombinations(t *testing.T) { + tests := []struct { + name string + oauthEnabled bool + transport string + enableMultiCluster bool + useLegacyTools bool + wantErr bool + }{ + { + name: "OAuth disabled with stdio", + oauthEnabled: false, + transport: "stdio", + enableMultiCluster: false, + useLegacyTools: false, + wantErr: false, + }, + { + name: "OAuth enabled with SSE", + oauthEnabled: true, + transport: "sse", + enableMultiCluster: false, + useLegacyTools: false, + wantErr: false, + }, + { + name: "OAuth enabled with streamable-http", + oauthEnabled: true, + transport: "streamable-http", + enableMultiCluster: false, + useLegacyTools: false, + wantErr: false, + }, + { + name: "Multi-cluster with unified tools", + oauthEnabled: false, + transport: "sse", + enableMultiCluster: true, + useLegacyTools: false, + wantErr: false, + }, + { + name: "Single cluster with legacy tools", + oauthEnabled: false, + transport: "stdio", + enableMultiCluster: false, + useLegacyTools: true, + wantErr: false, + }, + { + name: "All features compatible", + oauthEnabled: true, + transport: "sse", + enableMultiCluster: true, + useLegacyTools: false, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := NewConfig() + cfg.OAuthConfig.Enabled = tt.oauthEnabled + cfg.Transport = tt.transport + cfg.EnableMultiCluster = tt.enableMultiCluster + cfg.UseLegacyTools = tt.useLegacyTools + + err := cfg.ValidateConfig() + if (err != nil) != tt.wantErr { + t.Errorf("ValidateConfig() error = %v, wantErr %v", err, tt.wantErr) + } + }) + } +} + +func TestValidateConfig_InvalidCombinations(t *testing.T) { + tests := []struct { + name string + oauthEnabled bool + transport string + enableMultiCluster bool + useLegacyTools bool + expectedErrMsg string + }{ + { + name: "OAuth with stdio", + oauthEnabled: true, + transport: "stdio", + enableMultiCluster: false, + useLegacyTools: false, + expectedErrMsg: "OAuth authentication is not supported with stdio transport", + }, + { + name: "Multi-cluster with stdio", + oauthEnabled: false, + transport: "stdio", + enableMultiCluster: true, + useLegacyTools: false, + expectedErrMsg: "multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport", + }, + { + name: "Multi-cluster with legacy tools", + oauthEnabled: false, + transport: "sse", + enableMultiCluster: true, + useLegacyTools: true, + expectedErrMsg: "multi-cluster mode (--enable-multi-cluster) requires unified tools", + }, + { + name: "All invalid combinations", + oauthEnabled: true, + transport: "stdio", + enableMultiCluster: true, + useLegacyTools: true, + expectedErrMsg: "OAuth authentication is not supported with stdio transport", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + cfg := NewConfig() + cfg.OAuthConfig.Enabled = tt.oauthEnabled + cfg.Transport = tt.transport + cfg.EnableMultiCluster = tt.enableMultiCluster + cfg.UseLegacyTools = tt.useLegacyTools + + err := cfg.ValidateConfig() + if err == nil { + t.Fatal("Expected error, got nil") + } + + if !contains(err.Error(), tt.expectedErrMsg) { + t.Errorf("Expected error containing '%s', got '%s'", tt.expectedErrMsg, err.Error()) + } + }) + } +} diff --git a/internal/ctx/context.go b/internal/ctx/context.go index c7204015..374f4ac2 100644 --- a/internal/ctx/context.go +++ b/internal/ctx/context.go @@ -2,4 +2,7 @@ package ctx type ContextKey string +// AzureTokenKey is the context key for storing Azure tokens extracted from HTTP headers. +// This is the name of the HTTP header, not a hardcoded credential. +// #nosec G101 const AzureTokenKey ContextKey = "X-Azure-Token" diff --git a/internal/k8s/runcommand_executor_test.go b/internal/k8s/runcommand_executor_test.go new file mode 100644 index 00000000..d51891ca --- /dev/null +++ b/internal/k8s/runcommand_executor_test.go @@ -0,0 +1,281 @@ +package k8s + +import ( + "context" + "testing" + + "github.com/Azure/aks-mcp/internal/ctx" + k8sconfig "github.com/Azure/mcp-kubernetes/pkg/config" +) + +func TestExtractRequestContext_Success(t *testing.T) { + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "test-token-123") + + params := map[string]interface{}{ + "aks_resource_id": "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/cluster-test", + } + + reqCtx, err := extractRequestContext(c, params) + if err != nil { + t.Fatalf("Expected no error, got: %v", err) + } + + if reqCtx.AzureToken != "test-token-123" { + t.Errorf("Expected token 'test-token-123', got '%s'", reqCtx.AzureToken) + } + + if reqCtx.SubscriptionID != "sub-123" { + t.Errorf("Expected subscription 'sub-123', got '%s'", reqCtx.SubscriptionID) + } + + if reqCtx.ResourceGroup != "rg-test" { + t.Errorf("Expected resource group 'rg-test', got '%s'", reqCtx.ResourceGroup) + } + + if reqCtx.ClusterName != "cluster-test" { + t.Errorf("Expected cluster name 'cluster-test', got '%s'", reqCtx.ClusterName) + } +} + +func TestExtractRequestContext_MissingToken(t *testing.T) { + c := context.Background() + + params := map[string]interface{}{ + "aks_resource_id": "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/cluster-test", + } + + _, err := extractRequestContext(c, params) + if err == nil { + t.Fatal("Expected error for missing token, got nil") + } + + expectedMsg := "X-Azure-Token not found in context or empty" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestExtractRequestContext_EmptyToken(t *testing.T) { + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "") + + params := map[string]interface{}{ + "aks_resource_id": "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/cluster-test", + } + + _, err := extractRequestContext(c, params) + if err == nil { + t.Fatal("Expected error for empty token, got nil") + } + + expectedMsg := "X-Azure-Token not found in context or empty" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestExtractRequestContext_MissingResourceID(t *testing.T) { + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "test-token-123") + + params := map[string]interface{}{} + + _, err := extractRequestContext(c, params) + if err == nil { + t.Fatal("Expected error for missing resource ID, got nil") + } + + expectedMsg := "aks_resource_id not found in params or empty" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestExtractRequestContext_EmptyResourceID(t *testing.T) { + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "test-token-123") + + params := map[string]interface{}{ + "aks_resource_id": "", + } + + _, err := extractRequestContext(c, params) + if err == nil { + t.Fatal("Expected error for empty resource ID, got nil") + } + + expectedMsg := "aks_resource_id not found in params or empty" + if err.Error() != expectedMsg { + t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) + } +} + +func TestExtractRequestContext_InvalidResourceID(t *testing.T) { + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "test-token-123") + + testCases := []struct { + name string + resourceID string + expectError bool + }{ + { + name: "Invalid format", + resourceID: "invalid-resource-id", + expectError: true, + }, + { + name: "Missing cluster name", + resourceID: "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/", + expectError: true, + }, + { + name: "Wrong provider", + resourceID: "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.Compute/virtualMachines/vm-test", + expectError: false, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + params := map[string]interface{}{ + "aks_resource_id": tc.resourceID, + } + + _, err := extractRequestContext(c, params) + if tc.expectError && err == nil { + t.Fatal("Expected error, got nil") + } + if tc.expectError && err != nil { + t.Logf("Got expected error: %v", err) + } + }) + } +} + +func TestBuildCommand_ValidCommand(t *testing.T) { + executor := &RunCommandExecutor{} + + params := map[string]interface{}{ + "command": "kubectl get pods -n default", + } + + command, err := executor.buildCommand(params) + if err != nil { + t.Fatalf("Expected no error, got: %v", err) + } + + if command != "kubectl get pods -n default" { + t.Errorf("Expected command 'kubectl get pods -n default', got '%s'", command) + } +} + +func TestBuildCommand_MissingCommand(t *testing.T) { + executor := &RunCommandExecutor{} + + params := map[string]interface{}{} + + _, err := executor.buildCommand(params) + if err == nil { + t.Fatal("Expected error for missing command, got nil") + } +} + +func TestBuildCommand_EmptyCommand(t *testing.T) { + executor := &RunCommandExecutor{} + + params := map[string]interface{}{ + "command": "", + } + + command, err := executor.buildCommand(params) + if err != nil { + t.Fatalf("Expected no error for empty string, got: %v", err) + } + + if command != "" { + t.Errorf("Expected empty command, got '%s'", command) + } +} + +func TestBuildCommand_InvalidCommandType(t *testing.T) { + executor := &RunCommandExecutor{} + + params := map[string]interface{}{ + "command": 123, + } + + _, err := executor.buildCommand(params) + if err == nil { + t.Fatal("Expected error for invalid command type, got nil") + } +} + +func TestExecute_MissingContext(t *testing.T) { + executor := &RunCommandExecutor{} + + params := map[string]interface{}{ + "command": "kubectl get pods", + "aks_resource_id": "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/cluster-test", + } + + cfg := &k8sconfig.ConfigData{} + c := context.Background() + + _, err := executor.Execute(c, params, cfg) + if err == nil { + t.Fatal("Expected error for missing context, got nil") + } + + if err.Error() != "failed to extract request context: X-Azure-Token not found in context or empty" { + t.Errorf("Expected context error, got: %v", err) + } +} + +func TestExecute_InvalidParams(t *testing.T) { + executor := &RunCommandExecutor{} + + c := context.WithValue(context.Background(), ctx.AzureTokenKey, "test-token-123") + + testCases := []struct { + name string + params map[string]interface{} + expectedErr string + }{ + { + name: "Missing aks_resource_id", + params: map[string]interface{}{ + "command": "kubectl get pods", + }, + expectedErr: "failed to extract request context: aks_resource_id not found in params or empty", + }, + { + name: "Invalid aks_resource_id", + params: map[string]interface{}{ + "command": "kubectl get pods", + "aks_resource_id": "invalid", + }, + expectedErr: "failed to extract request context: failed to parse aks_resource_id", + }, + { + name: "Missing command", + params: map[string]interface{}{ + "aks_resource_id": "/subscriptions/sub-123/resourceGroups/rg-test/providers/Microsoft.ContainerService/managedClusters/cluster-test", + }, + expectedErr: "failed to build command: command parameter is required", + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + cfg := &k8sconfig.ConfigData{} + _, err := executor.Execute(c, tc.params, cfg) + if err == nil { + t.Fatal("Expected error, got nil") + } + if !containsError(err.Error(), tc.expectedErr) { + t.Errorf("Expected error containing '%s', got '%s'", tc.expectedErr, err.Error()) + } + }) + } +} + +func containsError(actual, expected string) bool { + return len(actual) >= len(expected) && actual[:len(expected)] == expected +} diff --git a/internal/server/context_test.go b/internal/server/context_test.go new file mode 100644 index 00000000..bf7c0129 --- /dev/null +++ b/internal/server/context_test.go @@ -0,0 +1,211 @@ +package server + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" + + "github.com/Azure/aks-mcp/internal/ctx" + "github.com/mark3labs/mcp-go/server" +) + +func TestSSEContextFunc_ExtractsToken(t *testing.T) { + cfg := createTestConfig("readonly", []string{}) + service := NewService(cfg) + err := service.Initialize() + if err != nil { + t.Fatalf("Failed to initialize service: %v", err) + } + + req := httptest.NewRequest("GET", "/sse", nil) + req.Header.Set("X-Azure-Token", "test-token-abc") + + var capturedContext context.Context + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + capturedContext = c + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if !ok { + t.Fatal("Expected token in context") + } + + if token != "test-token-abc" { + t.Errorf("Expected token 'test-token-abc', got '%s'", token) + } + + if capturedContext == nil { + t.Fatal("Context should have been captured") + } +} + +func TestSSEContextFunc_NoToken(t *testing.T) { + req := httptest.NewRequest("GET", "/sse", nil) + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if ok { + t.Errorf("Expected no token in context, got '%s'", token) + } +} + +func TestSSEContextFunc_EmptyToken(t *testing.T) { + req := httptest.NewRequest("GET", "/sse", nil) + req.Header.Set("X-Azure-Token", "") + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if ok { + t.Errorf("Expected no token in context for empty header, got '%s'", token) + } +} + +func TestStreamableHTTPContextFunc_ExtractsToken(t *testing.T) { + cfg := createTestConfig("readonly", []string{}) + service := NewService(cfg) + err := service.Initialize() + if err != nil { + t.Fatalf("Failed to initialize service: %v", err) + } + + req := httptest.NewRequest("POST", "/mcp", nil) + req.Header.Set("X-Azure-Token", "test-token-xyz") + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if !ok { + t.Fatal("Expected token in context") + } + + if token != "test-token-xyz" { + t.Errorf("Expected token 'test-token-xyz', got '%s'", token) + } +} + +func TestStreamableHTTPContextFunc_NoToken(t *testing.T) { + req := httptest.NewRequest("POST", "/mcp", nil) + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if ok { + t.Errorf("Expected no token in context, got '%s'", token) + } +} + +func TestMultipleHeaders_OnlyTokenExtracted(t *testing.T) { + req := httptest.NewRequest("POST", "/mcp", nil) + req.Header.Set("X-Azure-Token", "correct-token") + req.Header.Set("Authorization", "Bearer should-not-be-used") + req.Header.Set("X-Custom-Header", "custom-value") + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + c := contextFunc(context.Background(), req) + + token, ok := c.Value(ctx.AzureTokenKey).(string) + if !ok { + t.Fatal("Expected token in context") + } + + if token != "correct-token" { + t.Errorf("Expected token 'correct-token', got '%s'", token) + } +} + +func TestSSEServerCreation_WithContextFunc(t *testing.T) { + cfg := createTestConfig("readonly", []string{}) + service := NewService(cfg) + err := service.Initialize() + if err != nil { + t.Fatalf("Failed to initialize service: %v", err) + } + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + sseServer := server.NewSSEServer( + service.mcpServer, + server.WithSSEContextFunc(contextFunc), + ) + + if sseServer == nil { + t.Fatal("SSE server should not be nil") + } +} + +func TestStreamableHTTPServerCreation_WithContextFunc(t *testing.T) { + cfg := createTestConfig("readonly", []string{}) + service := NewService(cfg) + err := service.Initialize() + if err != nil { + t.Fatalf("Failed to initialize service: %v", err) + } + + contextFunc := func(c context.Context, r *http.Request) context.Context { + if token := r.Header.Get("X-Azure-Token"); token != "" { + c = context.WithValue(c, ctx.AzureTokenKey, token) + } + return c + } + + addr := "localhost:8080" + customServer := service.createCustomHTTPServerWithHelp404(addr) + + streamableServer := server.NewStreamableHTTPServer( + service.mcpServer, + server.WithStreamableHTTPServer(customServer), + server.WithHTTPContextFunc(contextFunc), + ) + + if streamableServer == nil { + t.Fatal("Streamable HTTP server should not be nil") + } +} From 73ff6c5cd23bf23c0916788ec52d08db9127dfbd Mon Sep 17 00:00:00 2001 From: Guoxun Wei Date: Fri, 23 Jan 2026 14:41:44 +0800 Subject: [PATCH 4/4] use flag --token-auth-only --- .../components/fleet/kubernetes/client.go | 2 +- internal/config/config.go | 50 ++--- internal/config/config_test.go | 188 +++++++++--------- internal/k8s/adapter.go | 8 +- internal/k8s/registry.go | 4 +- internal/server/server.go | 20 +- 6 files changed, 136 insertions(+), 136 deletions(-) diff --git a/internal/components/fleet/kubernetes/client.go b/internal/components/fleet/kubernetes/client.go index 1851f0f0..8559f08e 100644 --- a/internal/components/fleet/kubernetes/client.go +++ b/internal/components/fleet/kubernetes/client.go @@ -21,7 +21,7 @@ func NewClient() (*Client, error) { k8sExecutor := kubectl.NewExecutor() // Wrap it using the adapter to work with aks-mcp config - // Fleet operations don't use multi-cluster mode for now (always use local kubeconfig) + // Fleet operations don't support token-only authentication mode yet (always use local kubeconfig) wrappedExecutor := k8s.WrapK8sExecutor(k8sExecutor, false) return &Client{ diff --git a/internal/config/config.go b/internal/config/config.go index a529de8b..58316587 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -75,27 +75,27 @@ type ConfigData struct { // This flag is provided for backward compatibility and may be removed in future versions UseLegacyTools bool - // EnableMultiCluster enables multi-cluster mode for kubectl tools - // When enabled, kubectl commands are executed via Azure AKS RunCommand API - // When disabled (default), kubectl commands are executed locally via kubeconfig - EnableMultiCluster bool + // TokenAuthOnly enables token-only authentication mode for tools that support it + // When enabled, supported tools (e.g., kubectl) are executed via Azure AKS RunCommand API using user-provided tokens + // When disabled (default), tools are executed locally with default authentication (e.g., kubeconfig for Kubernetes tools) + TokenAuthOnly bool } // NewConfig creates and returns a new configuration instance func NewConfig() *ConfigData { return &ConfigData{ - Timeout: 60, - CacheTimeout: 1 * time.Minute, - SecurityConfig: security.NewSecurityConfig(), - OAuthConfig: auth.NewDefaultOAuthConfig(), - Transport: "stdio", - Port: 8000, - AccessLevel: "readonly", - EnabledComponents: []string{}, - AllowNamespaces: "", - LogLevel: "info", - UseLegacyTools: os.Getenv("USE_LEGACY_TOOLS") == "true", - EnableMultiCluster: false, + Timeout: 60, + CacheTimeout: 1 * time.Minute, + SecurityConfig: security.NewSecurityConfig(), + OAuthConfig: auth.NewDefaultOAuthConfig(), + Transport: "stdio", + Port: 8000, + AccessLevel: "readonly", + EnabledComponents: []string{}, + AllowNamespaces: "", + LogLevel: "info", + UseLegacyTools: os.Getenv("USE_LEGACY_TOOLS") == "true", + TokenAuthOnly: false, } } @@ -131,9 +131,9 @@ func (cfg *ConfigData) ParseFlags() { flag.StringVar(&cfg.AllowNamespaces, "allow-namespaces", "", "Comma-separated list of allowed Kubernetes namespaces (empty means all namespaces)") - // Multi-cluster configuration - flag.BoolVar(&cfg.EnableMultiCluster, "enable-multi-cluster", false, - "Enable multi-cluster mode for kubectl (uses Azure AKS RunCommand API instead of local kubeconfig)") + // Token-only authentication configuration + flag.BoolVar(&cfg.TokenAuthOnly, "token-auth-only", false, + "Enable token-only authentication mode for supported tools (e.g., kubectl uses Azure AKS RunCommand API with user-provided tokens instead of local kubeconfig)") // Logging settings flag.StringVar(&cfg.LogLevel, "log-level", "info", "Log level (debug, info, warn, error)") @@ -278,14 +278,14 @@ func (cfg *ConfigData) ValidateConfig() error { return fmt.Errorf("OAuth authentication is not supported with stdio transport per MCP specification") } - // Validate multi-cluster + stdio transport compatibility - if cfg.EnableMultiCluster && cfg.Transport == "stdio" { - return fmt.Errorf("multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport, use sse or streamable-http instead") + // Validate token-only authentication + stdio transport compatibility + if cfg.TokenAuthOnly && cfg.Transport == "stdio" { + return fmt.Errorf("token-only authentication mode (--token-auth-only) is not supported with stdio transport, use sse or streamable-http instead") } - // Validate multi-cluster + legacy tools compatibility - if cfg.EnableMultiCluster && cfg.UseLegacyTools { - return fmt.Errorf("multi-cluster mode (--enable-multi-cluster) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)") + // Validate token-only authentication + legacy tools compatibility + if cfg.TokenAuthOnly && cfg.UseLegacyTools { + return fmt.Errorf("token-only authentication mode (--token-auth-only) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)") } return nil diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 834e6ced..66a0cd2e 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -375,142 +375,142 @@ func TestValidateConfig_OAuthWithStreamableHTTP(t *testing.T) { } } -func TestValidateConfig_MultiClusterWithLegacyTools(t *testing.T) { +func TestValidateConfig_TokenAuthOnlyWithLegacyTools(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = true + cfg.TokenAuthOnly = true cfg.UseLegacyTools = true cfg.Transport = "sse" err := cfg.ValidateConfig() if err == nil { - t.Fatal("Expected error when multi-cluster is enabled with legacy tools, got nil") + t.Fatal("Expected error when token-only authentication is enabled with legacy tools, got nil") } - expectedMsg := "multi-cluster mode (--enable-multi-cluster) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)" + expectedMsg := "token-only authentication mode (--token-auth-only) requires unified tools and is not compatible with legacy tools (USE_LEGACY_TOOLS=true)" if err.Error() != expectedMsg { t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) } } -func TestValidateConfig_MultiClusterWithStdio(t *testing.T) { +func TestValidateConfig_TokenAuthOnlyWithStdio(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = true + cfg.TokenAuthOnly = true cfg.Transport = "stdio" err := cfg.ValidateConfig() if err == nil { - t.Fatal("Expected error when multi-cluster is enabled with stdio transport, got nil") + t.Fatal("Expected error when token-only authentication is enabled with stdio transport, got nil") } - expectedMsg := "multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport, use sse or streamable-http instead" + expectedMsg := "token-only authentication mode (--token-auth-only) is not supported with stdio transport, use sse or streamable-http instead" if err.Error() != expectedMsg { t.Errorf("Expected error '%s', got '%s'", expectedMsg, err.Error()) } } -func TestValidateConfig_MultiClusterWithSSE(t *testing.T) { +func TestValidateConfig_TokenAuthOnlyWithSSE(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = true + cfg.TokenAuthOnly = true cfg.Transport = "sse" cfg.UseLegacyTools = false err := cfg.ValidateConfig() if err != nil { - t.Errorf("Expected no error for multi-cluster with SSE transport, got: %v", err) + t.Errorf("Expected no error for token-only authentication with SSE transport, got: %v", err) } } -func TestValidateConfig_MultiClusterWithStreamableHTTP(t *testing.T) { +func TestValidateConfig_TokenAuthOnlyWithStreamableHTTP(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = true + cfg.TokenAuthOnly = true cfg.Transport = "streamable-http" cfg.UseLegacyTools = false err := cfg.ValidateConfig() if err != nil { - t.Errorf("Expected no error for multi-cluster with streamable-http transport, got: %v", err) + t.Errorf("Expected no error for token-only authentication with streamable-http transport, got: %v", err) } } -func TestValidateConfig_MultiClusterWithUnifiedTools(t *testing.T) { +func TestValidateConfig_TokenAuthOnlyWithUnifiedTools(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = true + cfg.TokenAuthOnly = true cfg.UseLegacyTools = false cfg.Transport = "sse" err := cfg.ValidateConfig() if err != nil { - t.Errorf("Expected no error for multi-cluster with unified tools, got: %v", err) + t.Errorf("Expected no error for token-only authentication with unified tools, got: %v", err) } } -func TestValidateConfig_LegacyToolsWithoutMultiCluster(t *testing.T) { +func TestValidateConfig_LegacyToolsWithoutTokenAuthOnly(t *testing.T) { cfg := NewConfig() - cfg.EnableMultiCluster = false + cfg.TokenAuthOnly = false cfg.UseLegacyTools = true err := cfg.ValidateConfig() if err != nil { - t.Errorf("Expected no error for legacy tools without multi-cluster, got: %v", err) + t.Errorf("Expected no error for legacy tools without token-only authentication, got: %v", err) } } func TestValidateConfig_ValidCombinations(t *testing.T) { tests := []struct { - name string - oauthEnabled bool - transport string - enableMultiCluster bool - useLegacyTools bool - wantErr bool + name string + oauthEnabled bool + transport string + tokenAuthOnly bool + useLegacyTools bool + wantErr bool }{ { - name: "OAuth disabled with stdio", - oauthEnabled: false, - transport: "stdio", - enableMultiCluster: false, - useLegacyTools: false, - wantErr: false, + name: "OAuth disabled with stdio", + oauthEnabled: false, + transport: "stdio", + tokenAuthOnly: false, + useLegacyTools: false, + wantErr: false, }, { - name: "OAuth enabled with SSE", - oauthEnabled: true, - transport: "sse", - enableMultiCluster: false, - useLegacyTools: false, - wantErr: false, + name: "OAuth enabled with SSE", + oauthEnabled: true, + transport: "sse", + tokenAuthOnly: false, + useLegacyTools: false, + wantErr: false, }, { - name: "OAuth enabled with streamable-http", - oauthEnabled: true, - transport: "streamable-http", - enableMultiCluster: false, - useLegacyTools: false, - wantErr: false, + name: "OAuth enabled with streamable-http", + oauthEnabled: true, + transport: "streamable-http", + tokenAuthOnly: false, + useLegacyTools: false, + wantErr: false, }, { - name: "Multi-cluster with unified tools", - oauthEnabled: false, - transport: "sse", - enableMultiCluster: true, - useLegacyTools: false, - wantErr: false, + name: "Token-only authentication with unified tools", + oauthEnabled: false, + transport: "sse", + tokenAuthOnly: true, + useLegacyTools: false, + wantErr: false, }, { - name: "Single cluster with legacy tools", - oauthEnabled: false, - transport: "stdio", - enableMultiCluster: false, - useLegacyTools: true, - wantErr: false, + name: "Single cluster with legacy tools", + oauthEnabled: false, + transport: "stdio", + tokenAuthOnly: false, + useLegacyTools: true, + wantErr: false, }, { - name: "All features compatible", - oauthEnabled: true, - transport: "sse", - enableMultiCluster: true, - useLegacyTools: false, - wantErr: false, + name: "All features compatible", + oauthEnabled: true, + transport: "sse", + tokenAuthOnly: true, + useLegacyTools: false, + wantErr: false, }, } @@ -519,7 +519,7 @@ func TestValidateConfig_ValidCombinations(t *testing.T) { cfg := NewConfig() cfg.OAuthConfig.Enabled = tt.oauthEnabled cfg.Transport = tt.transport - cfg.EnableMultiCluster = tt.enableMultiCluster + cfg.TokenAuthOnly = tt.tokenAuthOnly cfg.UseLegacyTools = tt.useLegacyTools err := cfg.ValidateConfig() @@ -532,44 +532,44 @@ func TestValidateConfig_ValidCombinations(t *testing.T) { func TestValidateConfig_InvalidCombinations(t *testing.T) { tests := []struct { - name string - oauthEnabled bool - transport string - enableMultiCluster bool - useLegacyTools bool - expectedErrMsg string + name string + oauthEnabled bool + transport string + tokenAuthOnly bool + useLegacyTools bool + expectedErrMsg string }{ { - name: "OAuth with stdio", - oauthEnabled: true, - transport: "stdio", - enableMultiCluster: false, - useLegacyTools: false, - expectedErrMsg: "OAuth authentication is not supported with stdio transport", + name: "OAuth with stdio", + oauthEnabled: true, + transport: "stdio", + tokenAuthOnly: false, + useLegacyTools: false, + expectedErrMsg: "OAuth authentication is not supported with stdio transport", }, { - name: "Multi-cluster with stdio", - oauthEnabled: false, - transport: "stdio", - enableMultiCluster: true, - useLegacyTools: false, - expectedErrMsg: "multi-cluster mode (--enable-multi-cluster) is not supported with stdio transport", + name: "Token-only authentication with stdio", + oauthEnabled: false, + transport: "stdio", + tokenAuthOnly: true, + useLegacyTools: false, + expectedErrMsg: "token-only authentication mode (--token-auth-only) is not supported with stdio transport", }, { - name: "Multi-cluster with legacy tools", - oauthEnabled: false, - transport: "sse", - enableMultiCluster: true, - useLegacyTools: true, - expectedErrMsg: "multi-cluster mode (--enable-multi-cluster) requires unified tools", + name: "Token-only authentication with legacy tools", + oauthEnabled: false, + transport: "sse", + tokenAuthOnly: true, + useLegacyTools: true, + expectedErrMsg: "token-only authentication mode (--token-auth-only) requires unified tools", }, { - name: "All invalid combinations", - oauthEnabled: true, - transport: "stdio", - enableMultiCluster: true, - useLegacyTools: true, - expectedErrMsg: "OAuth authentication is not supported with stdio transport", + name: "All invalid combinations", + oauthEnabled: true, + transport: "stdio", + tokenAuthOnly: true, + useLegacyTools: true, + expectedErrMsg: "OAuth authentication is not supported with stdio transport", }, } @@ -578,7 +578,7 @@ func TestValidateConfig_InvalidCombinations(t *testing.T) { cfg := NewConfig() cfg.OAuthConfig.Enabled = tt.oauthEnabled cfg.Transport = tt.transport - cfg.EnableMultiCluster = tt.enableMultiCluster + cfg.TokenAuthOnly = tt.tokenAuthOnly cfg.UseLegacyTools = tt.useLegacyTools err := cfg.ValidateConfig() diff --git a/internal/k8s/adapter.go b/internal/k8s/adapter.go index 9dae8828..851c96e8 100644 --- a/internal/k8s/adapter.go +++ b/internal/k8s/adapter.go @@ -60,11 +60,11 @@ func ConvertConfig(cfg *config.ConfigData) *k8sconfig.ConfigData { // WrapK8sExecutor makes an mcp-kubernetes CommandExecutor // compatible with the aks-mcp tools.CommandExecutor interface. -func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor, enableMultiCluster bool) tools.CommandExecutor { +func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor, tokenAuthOnly bool) tools.CommandExecutor { return &executorAdapter{ k8sExecutor: k8sExecutor, runCommandExecutor: NewRunCommandExecutor(), - enableMultiCluster: enableMultiCluster, + tokenAuthOnly: tokenAuthOnly, } } @@ -73,13 +73,13 @@ func WrapK8sExecutor(k8sExecutor k8stools.CommandExecutor, enableMultiCluster bo type executorAdapter struct { k8sExecutor k8stools.CommandExecutor runCommandExecutor *RunCommandExecutor - enableMultiCluster bool + tokenAuthOnly bool } // Execute adapts aks-mcp execution by converting its config // and delegating to the wrapped mcp-kubernetes executor or RunCommand executor. func (a *executorAdapter) Execute(ctx context.Context, params map[string]interface{}, cfg *config.ConfigData) (string, error) { - if a.enableMultiCluster { + if a.tokenAuthOnly { k8sCfg := ConvertConfig(cfg) return a.runCommandExecutor.Execute(ctx, params, k8sCfg) } diff --git a/internal/k8s/registry.go b/internal/k8s/registry.go index c1afcadc..317ac9ba 100644 --- a/internal/k8s/registry.go +++ b/internal/k8s/registry.go @@ -83,8 +83,8 @@ Examples: ) } -func RegisterKubectlTools(accessLevel string, useUnifiedTool bool, enableMultiCluster bool) []mcp.Tool { - if enableMultiCluster { +func RegisterKubectlTools(accessLevel string, useUnifiedTool bool, tokenAuthOnly bool) []mcp.Tool { + if tokenAuthOnly { return []mcp.Tool{ createCallKubectlTool(accessLevel), } diff --git a/internal/server/server.go b/internal/server/server.go index 28d4ea84..8a7dd00a 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -175,8 +175,8 @@ func (s *Service) registerAllComponents() { // Kubernetes Components s.registerKubernetesComponents() - if s.cfg.EnableMultiCluster { - logger.Infof("Multi-cluster mode enabled - skipping Azure component registration because they are not yet supported in multi-cluster mode") + if s.cfg.TokenAuthOnly { + logger.Infof("Token-only authentication mode enabled - skipping Azure component registration because they are not yet supported in token-only authentication mode") } else { // Azure Components s.registerAzureComponents() @@ -505,9 +505,9 @@ func (s *Service) registerKubernetesComponents() { // Core Kubernetes Component (kubectl) s.registerKubectlComponent() - // Do not register optional components in multi-cluster mode, they are not supported yet. - if s.cfg.EnableMultiCluster { - logger.Infof("Multi-cluster mode enabled - skipping optional Kubernetes component registration because they are not yet supported in multi-cluster mode") + // Do not register optional components in token-only authentication mode, they are not supported yet. + if s.cfg.TokenAuthOnly { + logger.Infof("Token-only authentication mode enabled - skipping optional Kubernetes component registration because they are not yet supported in token-only authentication mode") } else { // Optional Kubernetes Components (based on configuration) s.registerOptionalKubernetesComponents() @@ -529,15 +529,15 @@ func (s *Service) registerKubectlComponent() { } // Get kubectl tools filtered by access level and tool type - kubectlTools := k8s.RegisterKubectlTools(s.cfg.AccessLevel, useUnifiedTool, s.cfg.EnableMultiCluster) + kubectlTools := k8s.RegisterKubectlTools(s.cfg.AccessLevel, useUnifiedTool, s.cfg.TokenAuthOnly) // Create a kubectl executor kubectlExecutor := kubectl.NewKubectlToolExecutor() - // Wrap the executor with multi-cluster support if enabled - wrappedExecutor := k8s.WrapK8sExecutor(kubectlExecutor, s.cfg.EnableMultiCluster) - if s.cfg.EnableMultiCluster { - logger.Infof("Multi-cluster mode enabled: kubectl commands will use Azure AKS RunCommand API") + // Wrap the executor with token-only authentication support if enabled + wrappedExecutor := k8s.WrapK8sExecutor(kubectlExecutor, s.cfg.TokenAuthOnly) + if s.cfg.TokenAuthOnly { + logger.Infof("Token-only authentication mode enabled: supported tools will use Azure AKS RunCommand API with user-provided tokens") } // Register each kubectl tool