Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 41 additions & 12 deletions src/main/java/org/unlaxer/tinyexpression/mcp/McpServer.java
Original file line number Diff line number Diff line change
Expand Up @@ -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();
Expand All @@ -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);
Expand Down Expand Up @@ -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();

Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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;
Expand All @@ -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");
};
}

Expand Down Expand Up @@ -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 {
Expand Down
83 changes: 83 additions & 0 deletions src/test/java/org/unlaxer/tinyexpression/mcp/McpServerTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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();
Expand Down Expand Up @@ -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()
Expand Down Expand Up @@ -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<String> 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<String> httpGet(String path) throws Exception {
HttpRequest req = HttpRequest.newBuilder()
.uri(URI.create("http://127.0.0.1:" + port + path))
Expand Down
Loading