diff --git a/README.md b/README.md index 60ab9c4..10beae3 100644 --- a/README.md +++ b/README.md @@ -251,6 +251,39 @@ The reusable Spring Boot binding code lives in `FastMcpSafeProperties` and `FastMcpSafeConfigurationFactory`, but it is not a standalone Agent framework starter and does not create MCP clients by itself. +## Production integration checklist + +Production validation does not require FastMCP Java to implement the MCP +protocol. It validates that the safety wrapper remains the only model-facing +tool path when Spring AI, AgentScope, and the underlying MCP SDKs run against a +real MCP service. + +Before using a configured server in production: + +- Verify the model-facing tool list contains only virtual names such as + `get_my_orders`, never raw MCP names such as `getOrdersByUserId` or + `mcp__orders__getOrdersByUserId`. +- Verify virtual input schemas contain only model-fillable business arguments, + and do not expose protected arguments such as `userId`, `tenantId`, `role`, or + `includeDeleted`. +- Resolve protected values from server-side runtime context through resolver + beans; do not put sensitive values into `fastmcp.safe.*` configuration. +- For Spring AI production deployments, set + `fastmcp.safe.diagnostics.external-raw-provider=fail` unless the application + intentionally uses the documented external-provider compatibility path. +- Pass the safe provider, for example `fastMcpSafeToolCallbackProvider`, to the + model. Do not pass every `ToolCallbackProvider` bean as a collection unless raw + providers have been filtered out. +- Configure a `SafeAuditSink` and verify audits contain virtual/raw tool names, + caller/tenant identifiers, and injected argument names, but not injected + argument values. +- If the application streams tool events to AG-UI, CopilotKit, logs, or a + frontend, verify those events do not leak raw tool names, injected values, raw + arguments, or backend-only metadata. +- Run real MCP smoke tests from private deployment configuration only. Do not + commit real company domains, real MCP endpoints, credentials, or business tool + names to this repository. + ## Build ```bash diff --git a/README.zh-CN.md b/README.zh-CN.md index 9a0b2ad..aa1355a 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -239,6 +239,32 @@ starter 或应用自身提供 `Toolkit`;本库不负责创建 AgentScope agent `FastMcpSafeProperties` 和 `FastMcpSafeConfigurationFactory`,但它不是独立的 Agent 框架 starter,也不会自行创建 MCP client。 +## 生产接入检查清单 + +生产验证并不要求 FastMCP Java 自己实现 MCP 协议。它要验证的是:当 Spring AI、 +AgentScope 和底层 MCP SDK 连接真实 MCP 服务时,安全包装层仍然是唯一面向模型的 +tool 路径。 + +在生产使用某个配置 server 前,应至少核对: + +- 模型可见 tool 列表只包含 `get_my_orders` 这类 virtual names,不包含 + `getOrdersByUserId` 或 `mcp__orders__getOrdersByUserId` 这类 raw MCP names。 +- virtual input schema 只包含模型可填写的业务参数,不暴露 `userId`、`tenantId`、 + `role`、`includeDeleted` 等 protected arguments。 +- protected values 必须通过 resolver bean 从服务端运行时上下文解析,不写进 + `fastmcp.safe.*` 配置。 +- Spring AI 生产部署建议设置 + `fastmcp.safe.diagnostics.external-raw-provider=fail`,除非应用明确使用文档里的 + external-provider 兼容路径。 +- 模型侧只接收 safe provider,例如 `fastMcpSafeToolCallbackProvider`;不要把所有 + `ToolCallbackProvider` bean 作为集合直接交给模型,除非已经过滤掉 raw providers。 +- 配置 `SafeAuditSink`,并确认 audit 只包含 virtual/raw tool names、caller/tenant + 标识和 injected argument names,不包含 injected argument values。 +- 如果应用会把 tool events 流式发送到 AG-UI、CopilotKit、日志或前端,需要确认这些 + event 不泄露 raw tool names、注入值、raw arguments 或后端内部 metadata。 +- 真实 MCP smoke test 只应使用私有部署配置执行。不要把真实公司域名、真实 MCP + endpoint、凭据或业务 tool 名称提交到本仓库。 + ## 构建 ```bash diff --git a/fastmcp-examples/spring-ai-boot-starter/README.md b/fastmcp-examples/spring-ai-boot-starter/README.md index 7072b88..9963f2e 100644 --- a/fastmcp-examples/spring-ai-boot-starter/README.md +++ b/fastmcp-examples/spring-ai-boot-starter/README.md @@ -21,6 +21,13 @@ It covers: - a localhost fake MCP server exposing raw `searchCatalogByTenant`, while the model only sees the safe virtual `search_catalog(keyword)` callback +When adapting this example to an application, inject the safe provider named +`fastMcpSafeToolCallbackProvider` into the model wiring. Do not pass every +`ToolCallbackProvider` bean to the model unless raw providers have been filtered +out. For production Spring AI deployments, prefer +`fastmcp.safe.diagnostics.external-raw-provider=fail` so accidental external raw +providers fail startup instead of relying on logs. + Run it from the repository root with JDK 17 or newer: ```bash diff --git a/fastmcp-safe-core/src/main/java/io/github/sandking/fastmcp/safe/SafeMcpTool.java b/fastmcp-safe-core/src/main/java/io/github/sandking/fastmcp/safe/SafeMcpTool.java index 7fbc44f..9ea810a 100644 --- a/fastmcp-safe-core/src/main/java/io/github/sandking/fastmcp/safe/SafeMcpTool.java +++ b/fastmcp-safe-core/src/main/java/io/github/sandking/fastmcp/safe/SafeMcpTool.java @@ -54,8 +54,20 @@ public CompletionStage callAsync(Map input, SafeToolC recordAudit(safeContext, false, exception.code(), policyDecision); return failed(exception); } - CompletionStage rawResult = rawToolInvoker.callAsync(spec.rawServerName(), spec.rawToolName(), - Collections.unmodifiableMap(new LinkedHashMap<>(rawArguments)), safeContext); + CompletionStage rawResult; + try { + rawResult = rawToolInvoker.callAsync(spec.rawServerName(), spec.rawToolName(), + Collections.unmodifiableMap(new LinkedHashMap<>(rawArguments)), safeContext); + if (rawResult == null) { + SafeMcpException exception = new SafeMcpException("RAW_TOOL_FAILED", + "Raw tool invoker returned null"); + recordAudit(safeContext, false, exception.code(), "allow"); + return failed(exception); + } + } catch (RuntimeException exception) { + recordAudit(safeContext, false, "RAW_TOOL_FAILED", "allow"); + return failed(exception); + } return rawResult.handle((result, exception) -> { if (exception != null) { Throwable cause = unwrap(exception); diff --git a/fastmcp-safe-core/src/test/java/io/github/sandking/fastmcp/safe/SafeMcpToolTest.java b/fastmcp-safe-core/src/test/java/io/github/sandking/fastmcp/safe/SafeMcpToolTest.java index 722f942..8851313 100644 --- a/fastmcp-safe-core/src/test/java/io/github/sandking/fastmcp/safe/SafeMcpToolTest.java +++ b/fastmcp-safe-core/src/test/java/io/github/sandking/fastmcp/safe/SafeMcpToolTest.java @@ -69,6 +69,32 @@ void mapsVirtualArgumentsAndInjectsProtectedArguments() { assertEquals(SetSupport.of("userId", "tenantId"), auditEvents.get(0).injectedArgumentNames()); } + @Test + void synchronousRawInvokerFailureRecordsRawToolFailedAudit() { + List auditEvents = new ArrayList<>(); + RuntimeException rawFailure = new IllegalStateException("raw callback failed"); + SafeMcpTool tool = new SafeMcpTool("test", orderToolSpec(), + (serverName, rawToolName, rawArguments, context) -> { + throw rawFailure; + }, + SafeMcpPolicies.allow(), + auditEvents::add); + + CompletionException exception = assertThrows(CompletionException.class, + () -> tool.callAsync( + Map.of("status", "paid"), + SafeToolCallContext.builder().userId("user-1").tenantId("tenant-1").build()) + .toCompletableFuture() + .join()); + + assertEquals(rawFailure, exception.getCause()); + assertEquals(1, auditEvents.size()); + SafeAuditEvent event = auditEvents.get(0); + assertFalse(event.success()); + assertEquals("RAW_TOOL_FAILED", event.errorCode()); + assertEquals("allow", event.policyDecision()); + } + @Test void rejectsModelSuppliedProtectedRawArgument() { SafeMcpTool tool = new SafeMcpTool("test", orderToolSpec(), diff --git a/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/FastMcpSpringAiTools.java b/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/FastMcpSpringAiTools.java index 9e529db..d4b5af6 100644 --- a/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/FastMcpSpringAiTools.java +++ b/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/FastMcpSpringAiTools.java @@ -12,7 +12,6 @@ import io.github.sandking.fastmcp.safe.SafeMcpToolSpec; import io.github.sandking.fastmcp.safe.SafeToolCallContext; import io.github.sandking.fastmcp.safe.SafeToolResult; -import java.lang.reflect.Method; import java.util.ArrayList; import java.util.LinkedHashMap; import java.util.List; @@ -75,9 +74,10 @@ public static ToolCallbackProvider wrap(String defaultRawServerName, List rawTools = new LinkedHashMap<>(); for (ToolCallback rawCallback : rawCallbacks) { Objects.requireNonNull(rawCallback, "rawCallback must not be null"); - String rawServerName = originalServerName(rawCallback).orElse(defaultRawServerName); - putRawTool(rawTools, rawServerName, rawCallback.getToolDefinition().name(), rawCallback); - originalToolName(rawCallback).ifPresent(name -> putRawTool(rawTools, rawServerName, name, rawCallback)); + SpringAiRawToolIdentity identity = SpringAiRawToolIdentity.from(rawCallback, defaultRawServerName); + putRawTool(rawTools, identity.serverName(), identity.definitionToolName(), rawCallback); + identity.originalToolName().ifPresent(name -> + putRawTool(rawTools, identity.serverName(), name, rawCallback)); } List safeCallbacks = new ArrayList<>(); @@ -93,27 +93,6 @@ public static ToolCallbackProvider wrap(String defaultRawServerName, List originalToolName(ToolCallback rawCallback) { - return reflectedString(rawCallback, "getOriginalToolName"); - } - - private static java.util.Optional originalServerName(ToolCallback rawCallback) { - return reflectedString(rawCallback, "getOriginalServerName"); - } - - private static java.util.Optional reflectedString(ToolCallback rawCallback, String methodName) { - try { - Method method = rawCallback.getClass().getMethod(methodName); - Object value = method.invoke(rawCallback); - if (value instanceof String && !((String) value).trim().isEmpty()) { - return java.util.Optional.of((String) value); - } - return java.util.Optional.empty(); - } catch (ReflectiveOperationException exception) { - return java.util.Optional.empty(); - } - } - private static String rawKey(String rawServerName, String rawToolName) { return SpringAiMcpToolMapping.requireText(rawServerName, "rawServerName") + "\n" + SpringAiMcpToolMapping.requireText(rawToolName, "rawToolName"); diff --git a/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentity.java b/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentity.java new file mode 100644 index 0000000..d4e40c8 --- /dev/null +++ b/fastmcp-spring-ai-adapter/src/main/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentity.java @@ -0,0 +1,51 @@ +package io.github.sandking.fastmcp.springai; + +import java.lang.reflect.Method; +import java.util.Objects; +import java.util.Optional; +import org.springframework.ai.tool.ToolCallback; + +final class SpringAiRawToolIdentity { + private final String serverName; + private final String definitionToolName; + private final Optional originalToolName; + + private SpringAiRawToolIdentity(String serverName, String definitionToolName, Optional originalToolName) { + this.serverName = SpringAiMcpToolMapping.requireText(serverName, "serverName"); + this.definitionToolName = SpringAiMcpToolMapping.requireText(definitionToolName, "definitionToolName"); + this.originalToolName = Objects.requireNonNull(originalToolName, "originalToolName must not be null"); + } + + static SpringAiRawToolIdentity from(ToolCallback rawCallback, String defaultRawServerName) { + Objects.requireNonNull(rawCallback, "rawCallback must not be null"); + String serverName = reflectedString(rawCallback, "getOriginalServerName") + .orElse(SpringAiMcpToolMapping.requireText(defaultRawServerName, "defaultRawServerName")); + return new SpringAiRawToolIdentity(serverName, rawCallback.getToolDefinition().name(), + reflectedString(rawCallback, "getOriginalToolName")); + } + + String serverName() { + return serverName; + } + + String definitionToolName() { + return definitionToolName; + } + + Optional originalToolName() { + return originalToolName; + } + + private static Optional reflectedString(ToolCallback rawCallback, String methodName) { + try { + Method method = rawCallback.getClass().getMethod(methodName); + Object value = method.invoke(rawCallback); + if (value instanceof String && !((String) value).trim().isEmpty()) { + return Optional.of((String) value); + } + return Optional.empty(); + } catch (ReflectiveOperationException exception) { + return Optional.empty(); + } + } +} diff --git a/fastmcp-spring-ai-adapter/src/test/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentityTest.java b/fastmcp-spring-ai-adapter/src/test/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentityTest.java new file mode 100644 index 0000000..6827121 --- /dev/null +++ b/fastmcp-spring-ai-adapter/src/test/java/io/github/sandking/fastmcp/springai/SpringAiRawToolIdentityTest.java @@ -0,0 +1,114 @@ +package io.github.sandking.fastmcp.springai; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.model.ToolContext; +import org.springframework.ai.tool.ToolCallback; +import org.springframework.ai.tool.definition.DefaultToolDefinition; +import org.springframework.ai.tool.definition.ToolDefinition; + +class SpringAiRawToolIdentityTest { + @Test + void usesOriginalServerAndToolNamesWhenCallbacksExposeThem() { + SpringAiRawToolIdentity identity = SpringAiRawToolIdentity.from( + new OriginalNamesCallback("mcp__orders__getOrdersByUserId", "orders", "getOrdersByUserId"), + "spring-ai"); + + assertEquals("orders", identity.serverName()); + assertEquals("mcp__orders__getOrdersByUserId", identity.definitionToolName()); + assertTrue(identity.originalToolName().isPresent()); + assertEquals("getOrdersByUserId", identity.originalToolName().get()); + } + + @Test + void fallsBackToDefaultServerAndDefinitionNameWhenOriginalNamesAreUnavailable() { + SpringAiRawToolIdentity identity = SpringAiRawToolIdentity.from( + new BasicCallback("getOrdersByUserId"), + "spring-ai"); + + assertEquals("spring-ai", identity.serverName()); + assertEquals("getOrdersByUserId", identity.definitionToolName()); + assertFalse(identity.originalToolName().isPresent()); + } + + @Test + void ignoresBlankOriginalNames() { + SpringAiRawToolIdentity identity = SpringAiRawToolIdentity.from( + new OriginalNamesCallback("getOrdersByUserId", " ", ""), + "spring-ai"); + + assertEquals("spring-ai", identity.serverName()); + assertEquals("getOrdersByUserId", identity.definitionToolName()); + assertFalse(identity.originalToolName().isPresent()); + } + + @Test + void ignoresOriginalNameAccessorsThatThrow() { + SpringAiRawToolIdentity identity = SpringAiRawToolIdentity.from( + new ThrowingOriginalNamesCallback("getOrdersByUserId"), + "spring-ai"); + + assertEquals("spring-ai", identity.serverName()); + assertEquals("getOrdersByUserId", identity.definitionToolName()); + assertFalse(identity.originalToolName().isPresent()); + } + + private static class BasicCallback implements ToolCallback { + private final ToolDefinition toolDefinition; + + BasicCallback(String name) { + this.toolDefinition = new DefaultToolDefinition(name, "Raw tool", "{}"); + } + + @Override + public ToolDefinition getToolDefinition() { + return toolDefinition; + } + + @Override + public String call(String toolInput) { + return call(toolInput, new ToolContext(java.util.Map.of())); + } + + @Override + public String call(String toolInput, ToolContext toolContext) { + return "ok"; + } + } + + private static final class OriginalNamesCallback extends BasicCallback { + private final String originalServerName; + private final String originalToolName; + + private OriginalNamesCallback(String name, String originalServerName, String originalToolName) { + super(name); + this.originalServerName = originalServerName; + this.originalToolName = originalToolName; + } + + public String getOriginalServerName() { + return originalServerName; + } + + public String getOriginalToolName() { + return originalToolName; + } + } + + private static final class ThrowingOriginalNamesCallback extends BasicCallback { + private ThrowingOriginalNamesCallback(String name) { + super(name); + } + + public String getOriginalServerName() { + throw new IllegalStateException("server name unavailable"); + } + + public String getOriginalToolName() { + throw new IllegalStateException("tool name unavailable"); + } + } +} diff --git a/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/FastMcpSafeAutoConfiguration.java b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/FastMcpSafeAutoConfiguration.java index 01e1881..e39a850 100644 --- a/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/FastMcpSafeAutoConfiguration.java +++ b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/FastMcpSafeAutoConfiguration.java @@ -1,25 +1,16 @@ package io.github.sandking.fastmcp.springai.boot; -import io.github.sandking.fastmcp.safe.SafeAuditEvent; import io.github.sandking.fastmcp.safe.SafeAuditSink; import io.github.sandking.fastmcp.safe.boot.FastMcpSafeConfigurationFactory; import io.github.sandking.fastmcp.safe.boot.FastMcpSafeProperties; import io.github.sandking.fastmcp.safe.config.SafeMcpConfiguration; -import io.github.sandking.fastmcp.safe.config.SafeMcpServerConfiguration; -import io.github.sandking.fastmcp.springai.FastMcpSpringAiTools; -import io.github.sandking.fastmcp.springai.SpringAiMcpToolMapping; import io.github.sandking.fastmcp.springai.SpringAiToolArgumentResolver; import io.modelcontextprotocol.client.McpSyncClient; import java.util.ArrayList; -import java.util.Arrays; import java.util.LinkedHashMap; import java.util.List; -import java.util.Locale; import java.util.Map; import java.util.stream.Collectors; -import org.slf4j.Logger; -import org.slf4j.LoggerFactory; -import org.springframework.ai.tool.ToolCallback; import org.springframework.ai.tool.ToolCallbackProvider; import org.springframework.beans.factory.ListableBeanFactory; import org.springframework.beans.factory.ObjectProvider; @@ -37,13 +28,6 @@ @ConditionalOnProperty(prefix = "fastmcp.safe", name = "enabled", havingValue = "true", matchIfMissing = true) @EnableConfigurationProperties(FastMcpSafeProperties.class) public class FastMcpSafeAutoConfiguration { - private static final Logger logger = LoggerFactory.getLogger(FastMcpSafeAutoConfiguration.class); - private static final String EXTERNAL_RAW_PROVIDER_MESSAGE = - "External raw Spring AI ToolCallbackProvider beans are present: %s; ensure models receive " - + "fastMcpSafeToolCallbackProvider, not raw providers. This is a conservative diagnostic: " - + "external ToolCallbackProvider beans are the main raw-provider exposure risk, " - + "although they may include non-raw business providers."; - @Bean @ConditionalOnMissingBean public SafeMcpConfiguration fastMcpSafeConfiguration(FastMcpSafeProperties properties) { @@ -66,10 +50,9 @@ public ToolCallbackProvider fastMcpSafeToolCallbackProvider(SafeMcpConfiguration ObjectProvider auditSinkProvider, Map resolvers) { SafeAuditSink auditSink = auditSinkProvider.getIfAvailable(SafeAuditSink::noOp); - String externalRawProviderDiagnosticsMode = - normalizeExternalRawProviderDiagnosticsMode(properties.getDiagnostics().getExternalRawProvider()); List externalRawProviderNames = externalRawToolCallbackProviderNames(beanFactory); - diagnoseExternalRawProviders(externalRawProviderDiagnosticsMode, externalRawProviderNames, auditSink); + SpringAiExternalRawProviderDiagnostics.from(properties) + .diagnose(externalRawProviderNames, auditSink); List managedClients = managedClientFactory.createClients(configuration); @@ -81,25 +64,8 @@ public ToolCallbackProvider fastMcpSafeToolCallbackProvider(SafeMcpConfiguration } List externalRawProviders = externalRawToolCallbackProviders(beanFactory, externalRawProviderNames); - List externalRawCallbacks = externalRawProviders.stream() - .flatMap(provider -> Arrays.stream(provider.getToolCallbacks())) - .collect(Collectors.toList()); - List safeCallbacks = new ArrayList<>(); - for (SafeMcpServerConfiguration server : configuration.servers().values()) { - if (server.tools().isEmpty()) { - continue; - } - List rawCallbacks = managedRawProviders.containsKey(server.name()) - ? Arrays.asList(managedRawProviders.get(server.name()).getToolCallbacks()) - : externalRawCallbacks; - List mappings = server.tools().values().stream() - .map(tool -> SpringAiMcpToolMapping.from(server.name(), tool, resolvers)) - .collect(Collectors.toList()); - safeCallbacks.addAll(Arrays.asList( - FastMcpSpringAiTools.wrap(server.name(), rawCallbacks, mappings, auditSink) - .getToolCallbacks())); - } - ToolCallbackProvider safeProvider = ToolCallbackProvider.from(safeCallbacks); + ToolCallbackProvider safeProvider = new SpringAiSafeToolCallbackProviderAssembler().assemble( + configuration, managedRawProviders, externalRawProviders, resolvers, auditSink); List clients = managedClients.stream() .map(FastMcpSpringAiManagedClientFactory.ManagedMcpClient::client) .collect(Collectors.toList()); @@ -130,30 +96,6 @@ private List externalRawToolCallbackProviders(ListableBean return providers; } - private String normalizeExternalRawProviderDiagnosticsMode(String mode) { - String normalizedMode = mode == null ? "warn" : mode.trim().toLowerCase(Locale.ROOT); - if (!"warn".equals(normalizedMode) && !"fail".equals(normalizedMode) && !"off".equals(normalizedMode)) { - throw new IllegalArgumentException("Unsupported fastmcp.safe.diagnostics.external-raw-provider: " + mode); - } - return normalizedMode; - } - - private void diagnoseExternalRawProviders(String mode, List externalRawProviderNames, - SafeAuditSink auditSink) { - if (externalRawProviderNames.isEmpty() || "off".equals(mode)) { - return; - } - String message = String.format(EXTERNAL_RAW_PROVIDER_MESSAGE, externalRawProviderNames); - auditSink.record(SafeAuditEvent.diagnostic("spring-ai", - "EXTERNAL_RAW_PROVIDER_PRESENT", - "fastMcpSafeToolCallbackProvider", - Map.of("providerNames", String.join(",", externalRawProviderNames), "mode", mode))); - if ("fail".equals(mode)) { - throw new IllegalStateException(message); - } - logger.warn(message); - } - private void closeManagedClients(List managedClients, Throwable failure) { for (FastMcpSpringAiManagedClientFactory.ManagedMcpClient managedClient : managedClients) { diff --git a/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnostics.java b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnostics.java new file mode 100644 index 0000000..77d0778 --- /dev/null +++ b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnostics.java @@ -0,0 +1,54 @@ +package io.github.sandking.fastmcp.springai.boot; + +import io.github.sandking.fastmcp.safe.SafeAuditEvent; +import io.github.sandking.fastmcp.safe.SafeAuditSink; +import io.github.sandking.fastmcp.safe.boot.FastMcpSafeProperties; +import java.util.List; +import java.util.Locale; +import java.util.Map; +import org.slf4j.Logger; +import org.slf4j.LoggerFactory; + +final class SpringAiExternalRawProviderDiagnostics { + private static final Logger logger = LoggerFactory.getLogger(SpringAiExternalRawProviderDiagnostics.class); + private static final String MESSAGE = + "External raw Spring AI ToolCallbackProvider beans are present: %s; ensure models receive " + + "fastMcpSafeToolCallbackProvider, not raw providers. This is a conservative diagnostic: " + + "external ToolCallbackProvider beans are the main raw-provider exposure risk, " + + "although they may include non-raw business providers."; + + private final String mode; + + private SpringAiExternalRawProviderDiagnostics(String mode) { + this.mode = mode; + } + + static SpringAiExternalRawProviderDiagnostics from(FastMcpSafeProperties properties) { + return new SpringAiExternalRawProviderDiagnostics(normalize(properties.getDiagnostics() + .getExternalRawProvider())); + } + + void diagnose(List externalRawProviderNames, SafeAuditSink auditSink) { + if (externalRawProviderNames.isEmpty() || "off".equals(mode)) { + return; + } + SafeAuditSink sink = auditSink == null ? SafeAuditSink.noOp() : auditSink; + String message = String.format(MESSAGE, externalRawProviderNames); + sink.record(SafeAuditEvent.diagnostic("spring-ai", + "EXTERNAL_RAW_PROVIDER_PRESENT", + "fastMcpSafeToolCallbackProvider", + Map.of("providerNames", String.join(",", externalRawProviderNames), "mode", mode))); + if ("fail".equals(mode)) { + throw new IllegalStateException(message); + } + logger.warn(message); + } + + private static String normalize(String mode) { + String normalizedMode = mode == null ? "warn" : mode.trim().toLowerCase(Locale.ROOT); + if (!"warn".equals(normalizedMode) && !"fail".equals(normalizedMode) && !"off".equals(normalizedMode)) { + throw new IllegalArgumentException("Unsupported fastmcp.safe.diagnostics.external-raw-provider: " + mode); + } + return normalizedMode; + } +} diff --git a/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssembler.java b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssembler.java new file mode 100644 index 0000000..7c73ef2 --- /dev/null +++ b/fastmcp-spring-ai-boot-starter/src/main/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssembler.java @@ -0,0 +1,43 @@ +package io.github.sandking.fastmcp.springai.boot; + +import io.github.sandking.fastmcp.safe.SafeAuditSink; +import io.github.sandking.fastmcp.safe.config.SafeMcpConfiguration; +import io.github.sandking.fastmcp.safe.config.SafeMcpServerConfiguration; +import io.github.sandking.fastmcp.springai.FastMcpSpringAiTools; +import io.github.sandking.fastmcp.springai.SpringAiMcpToolMapping; +import io.github.sandking.fastmcp.springai.SpringAiToolArgumentResolver; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; +import org.springframework.ai.tool.ToolCallback; +import org.springframework.ai.tool.ToolCallbackProvider; + +final class SpringAiSafeToolCallbackProviderAssembler { + ToolCallbackProvider assemble(SafeMcpConfiguration configuration, + Map managedRawProviders, + List externalRawProviders, + Map resolvers, + SafeAuditSink auditSink) { + List externalRawCallbacks = externalRawProviders.stream() + .flatMap(provider -> Arrays.stream(provider.getToolCallbacks())) + .collect(Collectors.toList()); + List safeCallbacks = new ArrayList<>(); + for (SafeMcpServerConfiguration server : configuration.servers().values()) { + if (server.tools().isEmpty()) { + continue; + } + List rawCallbacks = managedRawProviders.containsKey(server.name()) + ? Arrays.asList(managedRawProviders.get(server.name()).getToolCallbacks()) + : externalRawCallbacks; + List mappings = server.tools().values().stream() + .map(tool -> SpringAiMcpToolMapping.from(server.name(), tool, resolvers)) + .collect(Collectors.toList()); + safeCallbacks.addAll(Arrays.asList( + FastMcpSpringAiTools.wrap(server.name(), rawCallbacks, mappings, auditSink) + .getToolCallbacks())); + } + return ToolCallbackProvider.from(safeCallbacks); + } +} diff --git a/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnosticsTest.java b/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnosticsTest.java new file mode 100644 index 0000000..ddef820 --- /dev/null +++ b/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiExternalRawProviderDiagnosticsTest.java @@ -0,0 +1,61 @@ +package io.github.sandking.fastmcp.springai.boot; + +import static org.assertj.core.api.Assertions.assertThat; +import static org.assertj.core.api.Assertions.assertThatThrownBy; + +import io.github.sandking.fastmcp.safe.SafeAuditEvent; +import io.github.sandking.fastmcp.safe.boot.FastMcpSafeProperties; +import java.util.ArrayList; +import java.util.List; +import org.junit.jupiter.api.Test; + +class SpringAiExternalRawProviderDiagnosticsTest { + @Test + void failModeRecordsDiagnosticAuditAndThrows() { + SpringAiExternalRawProviderDiagnostics diagnostics = SpringAiExternalRawProviderDiagnostics.from( + properties("fail")); + List events = new ArrayList<>(); + + assertThatThrownBy(() -> diagnostics.diagnose(List.of("rawOrderToolProvider"), events::add)) + .isInstanceOf(IllegalStateException.class) + .hasMessageContaining("External raw Spring AI ToolCallbackProvider beans are present") + .hasMessageContaining("rawOrderToolProvider"); + + assertThat(events).hasSize(1); + SafeAuditEvent event = events.get(0); + assertThat(event.eventType()).isEqualTo("DIAGNOSTIC"); + assertThat(event.framework()).isEqualTo("spring-ai"); + assertThat(event.virtualToolName()).isEqualTo("fastMcpSafeToolCallbackProvider"); + assertThat(event.errorCode()).isEqualTo("EXTERNAL_RAW_PROVIDER_PRESENT"); + assertThat(event.details()).containsEntry("providerNames", "rawOrderToolProvider") + .containsEntry("mode", "fail"); + } + + @Test + void offModeSkipsDiagnostics() { + SpringAiExternalRawProviderDiagnostics diagnostics = SpringAiExternalRawProviderDiagnostics.from( + properties("off")); + List events = new ArrayList<>(); + + diagnostics.diagnose(List.of("rawOrderToolProvider"), events::add); + + assertThat(events).isEmpty(); + } + + @Test + void blankOrUnknownModeFailsClearly() { + assertThatThrownBy(() -> SpringAiExternalRawProviderDiagnostics.from(properties(""))) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unsupported fastmcp.safe.diagnostics.external-raw-provider:"); + + assertThatThrownBy(() -> SpringAiExternalRawProviderDiagnostics.from(properties("block"))) + .isInstanceOf(IllegalArgumentException.class) + .hasMessageContaining("Unsupported fastmcp.safe.diagnostics.external-raw-provider: block"); + } + + private static FastMcpSafeProperties properties(String externalRawProviderMode) { + FastMcpSafeProperties properties = new FastMcpSafeProperties(); + properties.getDiagnostics().setExternalRawProvider(externalRawProviderMode); + return properties; + } +} diff --git a/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssemblerTest.java b/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssemblerTest.java new file mode 100644 index 0000000..66e504e --- /dev/null +++ b/fastmcp-spring-ai-boot-starter/src/test/java/io/github/sandking/fastmcp/springai/boot/SpringAiSafeToolCallbackProviderAssemblerTest.java @@ -0,0 +1,134 @@ +package io.github.sandking.fastmcp.springai.boot; + +import static org.assertj.core.api.Assertions.assertThat; + +import com.fasterxml.jackson.core.type.TypeReference; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.fasterxml.jackson.databind.node.JsonNodeFactory; +import com.fasterxml.jackson.databind.node.ObjectNode; +import io.github.sandking.fastmcp.safe.SafeAuditSink; +import io.github.sandking.fastmcp.safe.config.SafeMcpConfiguration; +import io.github.sandking.fastmcp.safe.config.SafeMcpServerConfiguration; +import io.github.sandking.fastmcp.safe.config.SafeMcpToolConfiguration; +import io.github.sandking.fastmcp.springai.SpringAiToolArgumentResolver; +import java.util.List; +import java.util.Map; +import java.util.concurrent.atomic.AtomicReference; +import org.junit.jupiter.api.Test; +import org.springframework.ai.chat.model.ToolContext; +import org.springframework.ai.tool.ToolCallback; +import org.springframework.ai.tool.ToolCallbackProvider; +import org.springframework.ai.tool.definition.DefaultToolDefinition; +import org.springframework.ai.tool.definition.ToolDefinition; + +class SpringAiSafeToolCallbackProviderAssemblerTest { + private static final ObjectMapper OBJECT_MAPPER = new ObjectMapper(); + private static final TypeReference> MAP_TYPE = new TypeReference<>() { + }; + + @Test + void usesManagedRawProviderForMatchingServerBeforeExternalProvider() { + CapturingToolCallback managedRawTool = new CapturingToolCallback("managed"); + CapturingToolCallback externalRawTool = new CapturingToolCallback("external"); + + ToolCallbackProvider safeProvider = new SpringAiSafeToolCallbackProviderAssembler().assemble( + configuration("orders"), + Map.of("orders", ToolCallbackProvider.from(managedRawTool)), + List.of(ToolCallbackProvider.from(externalRawTool)), + resolvers(), + SafeAuditSink.noOp()); + + String result = safeProvider.getToolCallbacks()[0].call("{\"status\":\"PAID\"}", + new ToolContext(Map.of("userId", "user-123"))); + + assertThat(result).isEqualTo("managed:user-123:PAID"); + assertThat(managedRawTool.lastInput()).containsEntry("orderStatus", "PAID") + .containsEntry("userId", "user-123"); + assertThat(externalRawTool.lastInput()).isNull(); + } + + @Test + void usesExternalRawProviderWhenServerHasNoManagedProvider() { + CapturingToolCallback externalRawTool = new CapturingToolCallback("external"); + + ToolCallbackProvider safeProvider = new SpringAiSafeToolCallbackProviderAssembler().assemble( + configuration("orders"), + Map.of(), + List.of(ToolCallbackProvider.from(externalRawTool)), + resolvers(), + SafeAuditSink.noOp()); + + String result = safeProvider.getToolCallbacks()[0].call("{\"status\":\"PAID\"}", + new ToolContext(Map.of("userId", "user-123"))); + + assertThat(result).isEqualTo("external:user-123:PAID"); + assertThat(externalRawTool.lastInput()).containsEntry("orderStatus", "PAID") + .containsEntry("userId", "user-123"); + } + + private static SafeMcpConfiguration configuration(String serverName) { + return SafeMcpConfiguration.builder() + .server(SafeMcpServerConfiguration.builder(serverName) + .tool(SafeMcpToolConfiguration.builder("getOrdersByUserId") + .name("get_my_orders") + .description("Get orders for the authenticated user.") + .inputSchema(virtualOrderSchema()) + .mapArgument("status", "orderStatus") + .injectArgument("userId", "currentUserId") + .build()) + .build()) + .build(); + } + + private static Map resolvers() { + return Map.of("currentUserId", context -> context.getContext().get("userId")); + } + + private static ObjectNode virtualOrderSchema() { + ObjectNode schema = JsonNodeFactory.instance.objectNode(); + schema.put("type", "object"); + ObjectNode properties = JsonNodeFactory.instance.objectNode(); + properties.set("status", JsonNodeFactory.instance.objectNode().put("type", "string")); + schema.set("properties", properties); + schema.putArray("required").add("status"); + return schema; + } + + private static final class CapturingToolCallback implements ToolCallback { + private final String source; + private final AtomicReference> lastInput = new AtomicReference<>(); + private final ToolDefinition toolDefinition = new DefaultToolDefinition("getOrdersByUserId", + "Raw order lookup", + "{\"type\":\"object\",\"properties\":{\"orderStatus\":{\"type\":\"string\"}," + + "\"userId\":{\"type\":\"string\"}},\"required\":[\"orderStatus\",\"userId\"]}"); + + private CapturingToolCallback(String source) { + this.source = source; + } + + @Override + public ToolDefinition getToolDefinition() { + return toolDefinition; + } + + @Override + public String call(String toolInput) { + return call(toolInput, new ToolContext(Map.of())); + } + + @Override + public String call(String toolInput, ToolContext toolContext) { + try { + Map input = OBJECT_MAPPER.readValue(toolInput, MAP_TYPE); + lastInput.set(input); + return source + ":" + input.get("userId") + ":" + input.get("orderStatus"); + } catch (Exception exception) { + throw new IllegalArgumentException(exception); + } + } + + Map lastInput() { + return lastInput.get(); + } + } +}