diff --git a/e2e/internal/testutil/mcp.go b/e2e/internal/testutil/mcp.go index 8e19cd27f..987747d5d 100644 --- a/e2e/internal/testutil/mcp.go +++ b/e2e/internal/testutil/mcp.go @@ -63,15 +63,49 @@ func (c *Client) CreateAPIKey(name string) string { return result.CreatePersonalAPIKey.Token } +// CreateOAuth2AccessToken creates a manual OAuth access token via the connect GraphQL API. +// It returns the raw bearer token string. +func (c *Client) CreateOAuth2AccessToken(name string, scopes []string) string { + const query = ` + mutation($input: CreateOAuth2AccessTokenInput!) { + createOAuth2AccessToken(input: $input) { + token + } + } + ` + + var result struct { + CreateOAuth2AccessToken struct { + Token string `json:"token"` + } `json:"createOAuth2AccessToken"` + } + + err := c.ExecuteConnect(query, map[string]any{ + "input": map[string]any{ + "name": name, + "expiresAt": time.Now().Add(90 * 24 * time.Hour).UTC().Format(time.RFC3339), + "scopes": scopes, + }, + }, &result) + require.NoError(c.T, err, "createOAuth2AccessToken failed") + require.NotEmpty(c.T, result.CreateOAuth2AccessToken.Token, "OAuth access token is empty") + + return result.CreateOAuth2AccessToken.Token +} + // NewMCPClient creates an MCP client authenticated with an API key. // It initializes an MCP session and stores the session ID. func NewMCPClient(t require.TestingT, owner *Client) *MCPClient { - token := owner.CreateAPIKey("e2e-mcp-test") + return NewMCPClientWithAccessToken(t, owner, owner.CreateAPIKey("e2e-mcp-test")) +} +// NewMCPClientWithAccessToken creates an MCP client authenticated with a bearer token. +// It initializes an MCP session and stores the session ID. +func NewMCPClientWithAccessToken(t require.TestingT, owner *Client, accessToken string) *MCPClient { mc := &MCPClient{ t: t, baseURL: owner.BaseURL() + "/mcp/v1", - apiToken: token, + apiToken: accessToken, httpClient: &http.Client{ Timeout: 30 * time.Second, }, diff --git a/e2e/mcp/oauth2_access_token_test.go b/e2e/mcp/oauth2_access_token_test.go new file mode 100644 index 000000000..911aa7d4c --- /dev/null +++ b/e2e/mcp/oauth2_access_token_test.go @@ -0,0 +1,57 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package mcp_test + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/e2e/internal/testutil" +) + +func TestMCP_OAuth2AccessToken_ListOrganizations(t *testing.T) { + t.Parallel() + + owner := testutil.NewClient(t, testutil.RoleOwner) + token := owner.CreateOAuth2AccessToken("e2e-mcp-oauth-token", []string{"v1:iam:read"}) + + mc := testutil.NewMCPClientWithAccessToken(t, owner, token) + + var result struct { + Organizations []struct { + ID string `json:"id"` + Name string `json:"name"` + } `json:"organizations"` + } + mc.CallToolInto("listOrganizations", map[string]any{}, &result) + + require.NotEmpty(t, result.Organizations) +} + +func TestMCP_OAuth2AccessToken_ScopeEnforcement(t *testing.T) { + t.Parallel() + + owner := testutil.NewClient(t, testutil.RoleOwner) + orgID := owner.GetOrganizationID().String() + token := owner.CreateOAuth2AccessToken("e2e-mcp-oauth-scope", []string{"v1:org:read"}) + + mc := testutil.NewMCPClientWithAccessToken(t, owner, token) + + msg := mc.CallToolExpectToolError("listThirdParties", map[string]any{ + "organizationId": orgID, + }) + assert.Equal(t, "insufficient scope", msg) +} diff --git a/pkg/server/api/mcp/v1/middleware.go b/pkg/server/api/mcp/v1/middleware.go deleted file mode 100644 index e43be765a..000000000 --- a/pkg/server/api/mcp/v1/middleware.go +++ /dev/null @@ -1,65 +0,0 @@ -// Copyright (c) 2025-2026 Probo Inc . -// -// Permission to use, copy, modify, and/or distribute this software for any -// purpose with or without fee is hereby granted, provided that the above -// copyright notice and this permission notice appear in all copies. -// -// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH -// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY -// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, -// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM -// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR -// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR -// PERFORMANCE OF THIS SOFTWARE. - -package mcp_v1 - -import ( - "errors" - "net/http" - - "go.gearno.de/kit/httpserver" - "go.gearno.de/kit/log" - "go.probo.inc/probo/pkg/server/api/authn" -) - -func RequireAPIKeyHandler( - logger *log.Logger, - next http.Handler, -) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - ctx := r.Context() - - correlationID := r.Header.Get("X-Request-ID") - if correlationID == "" { - correlationID = r.Header.Get("X-Correlation-ID") - } - - logger.InfoCtx( - ctx, - "MCP authentication attempt", - log.String("correlation_id", correlationID), - log.String("path", r.URL.Path), - ) - - apiKey := authn.APIKeyFromContext(ctx) - - identity := authn.IdentityFromContext(ctx) - if identity == nil { - w.Header().Set("WWW-Authenticate", "Bearer") - httpserver.RenderError(w, http.StatusUnauthorized, errors.New("authentication required")) - - return - } - - logger.InfoCtx( - ctx, - "MCP authentication successful", - log.String("correlation_id", correlationID), - log.String("identity_id", identity.ID.String()), - log.String("api_key_id", apiKey.ID.String()), - ) - - next.ServeHTTP(w, r.WithContext(ctx)) - }) -} diff --git a/pkg/server/api/mcp/v1/v1_handler.go b/pkg/server/api/mcp/v1/v1_handler.go index d1deb1510..69fd798df 100644 --- a/pkg/server/api/mcp/v1/v1_handler.go +++ b/pkg/server/api/mcp/v1/v1_handler.go @@ -82,7 +82,9 @@ func NewMux( r := chi.NewMux() r.Use(authn.NewAPIKeyMiddleware(iamSvc, tokenSecret)) - r.Handle("/", RequireAPIKeyHandler(logger, protectedHandler)) + r.Use(authn.NewOAuth2AccessTokenMiddleware(iamSvc)) + r.Use(authn.NewIdentityPresenceMiddleware()) + r.Handle("/", protectedHandler) logger.Info("MCP server initialized successfully")