Sanitize MCP errors to avoid leaking internal details

Signed-off-by: Sacha Al Himdani <sacha@getprobo.com>
This commit is contained in:
Sacha Al Himdani
2026-04-23 11:35:00 +02:00
parent 4f88b15f6d
commit d82df6b8ee
8 changed files with 152 additions and 35 deletions

View File

@@ -204,6 +204,20 @@ func (mc *MCPClient) CallTool(toolName string, args map[string]any) *MCPToolResu
return &toolResult
}
// CallToolExpectToolError invokes an MCP tool and expects a tool-level error
// (isError: true in the result). It returns the error text content.
func (mc *MCPClient) CallToolExpectToolError(toolName string, args map[string]any) string {
tr := mc.CallTool(toolName, args)
require.True(mc.t, tr.IsError, "expected tool %s to return isError", toolName)
require.NotEmpty(mc.t, tr.Content, "tool %s returned no content", toolName)
var text string
err := json.Unmarshal(tr.Content[0].Text, &text)
require.NoError(mc.t, err, "cannot unmarshal error text for %s", toolName)
return text
}
// CallToolInto invokes an MCP tool and unmarshals the first text content into dest.
func (mc *MCPClient) CallToolInto(toolName string, args map[string]any, dest any) {
tr := mc.CallTool(toolName, args)

View File

@@ -92,4 +92,27 @@ func TestMCP_Asset_CRUD(t *testing.T) {
"id": addResult.Asset.ID,
}, &deleteResult)
assert.Equal(t, addResult.Asset.ID, deleteResult.DeletedAssetID)
// Get deleted asset returns sanitized not-found error
msg := mc.CallToolExpectToolError("getAsset", map[string]any{
"id": addResult.Asset.ID,
})
assert.Equal(t, "resource not found", msg)
}
func TestMCP_Asset_PermissionDenied(t *testing.T) {
t.Parallel()
owner := testutil.NewClient(t, testutil.RoleOwner)
orgID := owner.GetOrganizationID().String()
viewer := testutil.NewClientInOrg(t, testutil.RoleViewer, owner)
viewerMC := testutil.NewMCPClient(t, viewer)
msg := viewerMC.CallToolExpectToolError("addAsset", map[string]any{
"organizationId": orgID,
"name": factory.SafeName("Asset"),
"amount": 1,
"assetType": "VIRTUAL",
"dataTypesStored": "PII",
})
assert.Contains(t, msg, "permission denied")
}

View File

@@ -86,4 +86,24 @@ func TestMCP_Risk_CRUD(t *testing.T) {
"id": addResult.Risk.ID,
}, &deleteResult)
assert.Equal(t, addResult.Risk.ID, deleteResult.DeletedRiskID)
// Get deleted risk returns sanitized not-found error
msg := mc.CallToolExpectToolError("getRisk", map[string]any{
"id": addResult.Risk.ID,
})
assert.Equal(t, "resource not found", msg)
}
func TestMCP_Risk_PermissionDenied(t *testing.T) {
t.Parallel()
owner := testutil.NewClient(t, testutil.RoleOwner)
orgID := owner.GetOrganizationID().String()
viewer := testutil.NewClientInOrg(t, testutil.RoleViewer, owner)
viewerMC := testutil.NewMCPClient(t, viewer)
msg := viewerMC.CallToolExpectToolError("addRisk", map[string]any{
"organizationId": orgID,
"name": factory.SafeName("Risk"),
})
assert.Contains(t, msg, "permission denied")
}

View File

@@ -76,4 +76,42 @@ func TestMCP_Vendor_CRUD(t *testing.T) {
"id": addResult.Vendor.ID,
}, &deleteResult)
assert.Equal(t, addResult.Vendor.ID, deleteResult.DeletedVendorID)
// Update deleted vendor returns sanitized not-found error
msg := mc.CallToolExpectToolError("updateVendor", map[string]any{
"id": addResult.Vendor.ID,
"name": "Should Fail",
})
assert.Equal(t, "resource not found", msg)
}
func TestMCP_Vendor_ValidationError(t *testing.T) {
t.Parallel()
owner := testutil.NewClient(t, testutil.RoleOwner)
mc := testutil.NewMCPClient(t, owner)
orgID := owner.GetOrganizationID().String()
msg := mc.CallToolExpectToolError("addVendor", map[string]any{
"organizationId": orgID,
"name": "",
})
assert.Contains(t, msg, "name")
assert.NotContains(t, msg, "pq:")
assert.NotContains(t, msg, "sql:")
}
func TestMCP_Vendor_PermissionDenied(t *testing.T) {
t.Parallel()
owner := testutil.NewClient(t, testutil.RoleOwner)
orgID := owner.GetOrganizationID().String()
viewer := testutil.NewClientInOrg(t, testutil.RoleViewer, owner)
viewerMC := testutil.NewMCPClient(t, viewer)
msg := viewerMC.CallToolExpectToolError("addVendor", map[string]any{
"organizationId": orgID,
"name": factory.SafeName("Vendor"),
})
assert.Contains(t, msg, "permission denied")
assert.NotContains(t, msg, "pq:")
assert.NotContains(t, msg, "sql:")
}