From a55c6e27b6c0043b54c579d072d81db4789a796d Mon Sep 17 00:00:00 2001 From: opaopa6969 Date: Tue, 1 Sep 2026 08:08:39 +0900 Subject: [PATCH] Fail fast on invalid MCP options (#96) --- .../unlaxer/tinyexpression/mcp/McpServer.java | 53 +++++++++--- .../tinyexpression/mcp/McpServerTest.java | 83 +++++++++++++++++++ 2 files changed, 124 insertions(+), 12 deletions(-) diff --git a/src/main/java/org/unlaxer/tinyexpression/mcp/McpServer.java b/src/main/java/org/unlaxer/tinyexpression/mcp/McpServer.java index 9340f5d5..322e0393 100644 --- a/src/main/java/org/unlaxer/tinyexpression/mcp/McpServer.java +++ b/src/main/java/org/unlaxer/tinyexpression/mcp/McpServer.java @@ -205,10 +205,12 @@ private ObjectNode handleToolsCall(JsonNode params) throws Exception { case "execute_batch" -> handleExecuteBatch(args); case "parity_check" -> handleParityCheck(args); case "list_backends" -> handleListBackends(); - default -> throw new IllegalArgumentException("Unknown tool: " + name); + default -> throw new UnknownToolException("Unknown tool: " + name); }; - } catch (IllegalArgumentException e) { + } catch (UnknownToolException e) { return errorResp(null, -32601, e.getMessage()); + } catch (IllegalArgumentException e) { + return errorResp(null, -32602, e.getMessage()); } ArrayNode content = MAPPER.createArrayNode(); @@ -225,12 +227,11 @@ private ObjectNode handleToolsCall(JsonNode params) throws Exception { private String handleEvaluate(JsonNode args) throws Exception { String formula = args.path("formula").asText(); - String backendStr = args.path("backend").asText("AST_EVALUATOR"); - String resultTypeStr = args.path("resultType").asText("float"); + String backendStr = optionText(args, "backend", "AST_EVALUATOR"); + String resultTypeStr = optionText(args, "resultType", "float"); JsonNode variablesNode = args.path("variables"); - ExecutionBackend backend = ExecutionBackend.parse(backendStr) - .orElse(ExecutionBackend.AST_EVALUATOR); + ExecutionBackend backend = parseBackend(backendStr); checkBackendAllowed(backend); ExpressionTypes resultType = parseResultType(resultTypeStr); @@ -268,7 +269,8 @@ private String handleEvaluate(JsonNode args) throws Exception { private String handleValidate(JsonNode args) throws Exception { String formula = args.path("formula").asText(); - String resultTypeStr = args.path("resultType").asText("float"); + String resultTypeStr = optionText(args, "resultType", "float"); + parseResultType(resultTypeStr); ObjectNode out = MAPPER.createObjectNode(); @@ -328,9 +330,9 @@ private String handleExecuteBatch(JsonNode args) throws Exception { bf.dependsOn.add(d.asText()); } } - String backendStr = fn.path("backend").asText("AST_EVALUATOR"); - String resultTypeStr = fn.path("resultType").asText("float"); - bf.backend = ExecutionBackend.parse(backendStr).orElse(ExecutionBackend.AST_EVALUATOR); + String backendStr = optionText(fn, "backend", "AST_EVALUATOR"); + String resultTypeStr = optionText(fn, "resultType", "float"); + bf.backend = parseBackend(backendStr); bf.resultType = parseResultType(resultTypeStr); checkBackendAllowed(bf.backend); batchFormulas.add(bf); @@ -417,7 +419,7 @@ private String handleExecuteBatch(JsonNode args) throws Exception { private String handleParityCheck(JsonNode args) throws Exception { String formula = args.path("formula").asText(); - String resultTypeStr = args.path("resultType").asText("float"); + String resultTypeStr = optionText(args, "resultType", "float"); JsonNode variablesNode = args.path("variables"); ExpressionTypes resultType = parseResultType(resultTypeStr); @@ -765,6 +767,25 @@ private static boolean isJavaCodeBackend(ExecutionBackend backend) { || backend == ExecutionBackend.JAVA_CODE_LEGACY_ASTCREATOR; } + private static String optionText(JsonNode options, String field, String defaultValue) { + if (options == null || !options.has(field)) { + return defaultValue; + } + JsonNode value = options.get(field); + if (value == null || !value.isTextual() || value.asText().isBlank()) { + throw new IllegalArgumentException( + "Option " + field + " must be a non-blank string when specified"); + } + return value.asText(); + } + + private static ExecutionBackend parseBackend(String backend) { + return ExecutionBackend.parse(backend).orElseThrow(() -> + new IllegalArgumentException( + "Unknown backend: " + backend + ". Allowed: " + + java.util.Arrays.toString(ExecutionBackend.values()))); + } + private static ExpressionTypes parseResultType(String resultTypeStr) { return switch (resultTypeStr.toLowerCase().strip()) { case "float", "_float", "number" -> ExpressionTypes._float; @@ -775,7 +796,9 @@ private static ExpressionTypes parseResultType(String resultTypeStr) { case "boolean", "_boolean", "bool" -> ExpressionTypes._boolean; case "object" -> ExpressionTypes.object; case "bigdecimal", "_bigdecimal" -> ExpressionTypes.bigDecimal; - default -> ExpressionTypes._float; + default -> throw new IllegalArgumentException( + "Unknown resultType: " + resultTypeStr + + ". Allowed: float, double, int, long, string, boolean, object, bigdecimal"); }; } @@ -864,6 +887,12 @@ private static class BatchFormula { ExpressionTypes resultType; } + private static final class UnknownToolException extends IllegalArgumentException { + private UnknownToolException(String message) { + super(message); + } + } + // ─── HTTP helpers ────────────────────────────────────────── private static String readBody(HttpExchange ex) throws IOException { diff --git a/src/test/java/org/unlaxer/tinyexpression/mcp/McpServerTest.java b/src/test/java/org/unlaxer/tinyexpression/mcp/McpServerTest.java index fb75a7d1..edbb93d4 100644 --- a/src/test/java/org/unlaxer/tinyexpression/mcp/McpServerTest.java +++ b/src/test/java/org/unlaxer/tinyexpression/mcp/McpServerTest.java @@ -85,6 +85,32 @@ public void toolsCall_evaluate_simple() throws Exception { assertEquals("AST_EVALUATOR", eval.get("backend_used").asText()); } + @Test + public void toolsCall_evaluate_unknownBackendIsInvalidParams() throws Exception { + ObjectNode args = MAPPER.createObjectNode(); + args.put("formula", "1+2"); + args.put("backend", "P4_MAGIC"); + + JsonNode error = callToolError("evaluate", args); + + assertEquals(-32602, error.get("code").asInt()); + assertTrue(error.get("message").asText().contains("P4_MAGIC")); + assertTrue(error.get("message").asText().contains("P4_AST_EVALUATOR")); + } + + @Test + public void toolsCall_evaluate_unknownResultTypeIsInvalidParams() throws Exception { + ObjectNode args = MAPPER.createObjectNode(); + args.put("formula", "1+2"); + args.put("resultType", "decimal128"); + + JsonNode error = callToolError("evaluate", args); + + assertEquals(-32602, error.get("code").asInt()); + assertTrue(error.get("message").asText().contains("decimal128")); + assertTrue(error.get("message").asText().contains("bigdecimal")); + } + @Test public void toolsCall_evaluate_withVariables() throws Exception { ObjectNode args = MAPPER.createObjectNode(); @@ -165,16 +191,52 @@ public void toolsCall_execute_batch_withDependency() throws Exception { for (JsonNode r : results) { if ("base".equals(r.get("name").asText())) { assertEquals(10.0, r.get("result").asDouble(), 0.001); + assertEquals("AST_EVALUATOR", r.get("backend_used").asText()); foundBase = true; } if ("total".equals(r.get("name").asText())) { assertEquals(110.0, r.get("result").asDouble(), 0.001); + assertEquals("AST_EVALUATOR", r.get("backend_used").asText()); foundTotal = true; } } assertTrue(foundBase && foundTotal); } + @Test + public void toolsCall_executeBatch_unknownOptionsFailBeforeExecution() throws Exception { + ObjectNode unknownBackend = MAPPER.createObjectNode(); + unknownBackend.put("name", "badBackend"); + unknownBackend.put("formula", "1+2"); + unknownBackend.put("backend", "P4_MAGIC"); + ObjectNode unknownResultType = MAPPER.createObjectNode(); + unknownResultType.put("name", "badType"); + unknownResultType.put("formula", "3+4"); + unknownResultType.put("resultType", "decimal128"); + + ObjectNode backendArgs = MAPPER.createObjectNode(); + backendArgs.set("formulas", MAPPER.createArrayNode().add(unknownBackend)); + JsonNode backendError = callToolError("execute_batch", backendArgs); + assertEquals(-32602, backendError.get("code").asInt()); + assertTrue(backendError.get("message").asText().contains("P4_MAGIC")); + + ObjectNode typeArgs = MAPPER.createObjectNode(); + typeArgs.set("formulas", MAPPER.createArrayNode().add(unknownResultType)); + JsonNode typeError = callToolError("execute_batch", typeArgs); + assertEquals(-32602, typeError.get("code").asInt()); + assertTrue(typeError.get("message").asText().contains("decimal128")); + } + + @Test + public void toolsCall_validateAndParity_unknownResultTypeIsInvalidParams() throws Exception { + ObjectNode args = MAPPER.createObjectNode(); + args.put("formula", "1+2"); + args.put("resultType", "decimal128"); + + assertEquals(-32602, callToolError("validate", args).get("code").asInt()); + assertEquals(-32602, callToolError("parity_check", args).get("code").asInt()); + } + @Test public void toolsCall_list_backends() throws Exception { JsonNode result = rpc("tools/call", MAPPER.createObjectNode() @@ -307,6 +369,27 @@ private JsonNode rpc(String method, JsonNode params) throws Exception { return respBody.get("result"); } + private JsonNode callToolError(String toolName, JsonNode arguments) throws Exception { + var reqBody = MAPPER.createObjectNode(); + reqBody.put("jsonrpc", "2.0"); + reqBody.put("id", 1); + reqBody.put("method", "tools/call"); + reqBody.set("params", MAPPER.createObjectNode() + .put("name", toolName) + .set("arguments", arguments)); + + HttpRequest req = HttpRequest.newBuilder() + .uri(URI.create("http://127.0.0.1:" + port + "/mcp")) + .header("Content-Type", "application/json") + .POST(HttpRequest.BodyPublishers.ofString(MAPPER.writeValueAsString(reqBody))) + .build(); + HttpResponse resp = client.send(req, HttpResponse.BodyHandlers.ofString()); + assertEquals(200, resp.statusCode()); + JsonNode response = MAPPER.readTree(resp.body()); + assertNotNull("Response: " + resp.body(), response.get("error")); + return response.get("error"); + } + private HttpResponse httpGet(String path) throws Exception { HttpRequest req = HttpRequest.newBuilder() .uri(URI.create("http://127.0.0.1:" + port + path))