From 73f392e271dd0605af7ab2241b0ab322f5df7a14 Mon Sep 17 00:00:00 2001 From: Lighthousexx <57421816+Lighthousexx@users.noreply.github.com> Date: Thu, 14 Aug 2025 20:05:45 +0800 Subject: [PATCH] feat: chatCompletion supports mcp tool (#37) --- .../ai/z/openapi/service/model/ChatTool.java | 3 + .../z/openapi/service/model/ChatToolType.java | 4 +- .../ai/z/openapi/service/model/MCPTool.java | 43 ++++++++++ .../z/openapi/service/model/MCPToolCall.java | 39 +++++++++ .../service/model/MCPToolDefinition.java | 50 ++++++++++++ .../service/model/McpToolTransportType.java | 16 ++++ .../ai/z/openapi/service/model/ToolCalls.java | 5 +- .../openapi/service/chat/ChatServiceTest.java | 81 +++++++++++++++++++ 8 files changed, 239 insertions(+), 2 deletions(-) create mode 100644 core/src/main/java/ai/z/openapi/service/model/MCPTool.java create mode 100644 core/src/main/java/ai/z/openapi/service/model/MCPToolCall.java create mode 100644 core/src/main/java/ai/z/openapi/service/model/MCPToolDefinition.java create mode 100644 core/src/main/java/ai/z/openapi/service/model/McpToolTransportType.java diff --git a/core/src/main/java/ai/z/openapi/service/model/ChatTool.java b/core/src/main/java/ai/z/openapi/service/model/ChatTool.java index fbfd654..a161b2d 100644 --- a/core/src/main/java/ai/z/openapi/service/model/ChatTool.java +++ b/core/src/main/java/ai/z/openapi/service/model/ChatTool.java @@ -23,4 +23,7 @@ public class ChatTool { @JsonProperty("web_search") private WebSearch webSearch; + @JsonProperty("mcp") + private MCPTool mcp; + } diff --git a/core/src/main/java/ai/z/openapi/service/model/ChatToolType.java b/core/src/main/java/ai/z/openapi/service/model/ChatToolType.java index 4ce0692..6278132 100644 --- a/core/src/main/java/ai/z/openapi/service/model/ChatToolType.java +++ b/core/src/main/java/ai/z/openapi/service/model/ChatToolType.java @@ -6,7 +6,9 @@ public enum ChatToolType { RETRIEVAL("retrieval"), - FUNCTION("function"); + FUNCTION("function"), + + MCP("mcp"),; private final String value; diff --git a/core/src/main/java/ai/z/openapi/service/model/MCPTool.java b/core/src/main/java/ai/z/openapi/service/model/MCPTool.java new file mode 100644 index 0000000..5b1d4b6 --- /dev/null +++ b/core/src/main/java/ai/z/openapi/service/model/MCPTool.java @@ -0,0 +1,43 @@ +package ai.z.openapi.service.model; + +import com.fasterxml.jackson.databind.JsonNode; +import com.fasterxml.jackson.databind.node.JsonNodeFactory; +import com.fasterxml.jackson.databind.node.ObjectNode; +import lombok.*; + +import java.util.Map; +import java.util.Set; + +@Data +@Builder +@NoArgsConstructor +@AllArgsConstructor +public class MCPTool { + + /** + * Identifier for the MCP server, used to distinguish different MCP servers, required + */ + private String server_label; + + /** + * URL of the MCP server, optional Default (if this field is empty): use server_label + * as mcpCode to connect to Zhipu AI's MCP servers + */ + private String server_url; + + /** + * Transport method for MCP calls: sse/streamable-http, defaults to streamable-http + */ + private String transport_type; + + /** + * List of allowed tools to call, defaults to empty (allowing all tools) + */ + private Set allowed_tools; + + /** + * Headers for connecting to MCP server, used for authentication + */ + private Map headers; + +} diff --git a/core/src/main/java/ai/z/openapi/service/model/MCPToolCall.java b/core/src/main/java/ai/z/openapi/service/model/MCPToolCall.java new file mode 100644 index 0000000..ebeb3c5 --- /dev/null +++ b/core/src/main/java/ai/z/openapi/service/model/MCPToolCall.java @@ -0,0 +1,39 @@ +package ai.z.openapi.service.model; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.io.Serializable; +import java.util.List; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class MCPToolCall implements Serializable { + + private static final long serialVersionUID = 2080214859980927710L; + + // Unique identifier for MCP tool call + private String id; + + private String type; + + private String server_label; + + private String error; + + // type = mcp_list_tools + private List tools; + + // type = mcp_call + // Tool call parameters, parameters as JSON string + private String arguments; + + // Tool name + private String name; + + // Tool result output + private Object output; + +} \ No newline at end of file diff --git a/core/src/main/java/ai/z/openapi/service/model/MCPToolDefinition.java b/core/src/main/java/ai/z/openapi/service/model/MCPToolDefinition.java new file mode 100644 index 0000000..0fcf807 --- /dev/null +++ b/core/src/main/java/ai/z/openapi/service/model/MCPToolDefinition.java @@ -0,0 +1,50 @@ +package ai.z.openapi.service.model; + +import lombok.AllArgsConstructor; +import lombok.Data; +import lombok.NoArgsConstructor; + +import java.io.Serializable; +import java.util.List; +import java.util.Map; + +@Data +@NoArgsConstructor +@AllArgsConstructor +public class MCPToolDefinition implements Serializable { + + private static final long serialVersionUID = -3960033025319205212L; + + // Tool name + private String name; + + // Tool description + private String description; + + private Object annotations; + + // Tool input parameter specification + private InputSchema input_schema; + + @Data + @NoArgsConstructor + @AllArgsConstructor + public static class InputSchema implements Serializable { + + private static final long serialVersionUID = 4723599134723995986L; + + // Fixed value "object" + private String type; + + // Parameter properties definition + private Map properties; + + // List of required properties + private List required; + + // Whether additional properties are allowed + private Boolean additionalProperties; + + } + +} diff --git a/core/src/main/java/ai/z/openapi/service/model/McpToolTransportType.java b/core/src/main/java/ai/z/openapi/service/model/McpToolTransportType.java new file mode 100644 index 0000000..66c1352 --- /dev/null +++ b/core/src/main/java/ai/z/openapi/service/model/McpToolTransportType.java @@ -0,0 +1,16 @@ +package ai.z.openapi.service.model; + +import lombok.AllArgsConstructor; +import lombok.Getter; + +@AllArgsConstructor +@Getter +public enum McpToolTransportType { + + SSE("sse", "SSE"), STREAMABLE_HTTP("streamable-http", "streamable http"); + + private final String code; + + private final String value; + +} diff --git a/core/src/main/java/ai/z/openapi/service/model/ToolCalls.java b/core/src/main/java/ai/z/openapi/service/model/ToolCalls.java index 484ae21..02a7edc 100644 --- a/core/src/main/java/ai/z/openapi/service/model/ToolCalls.java +++ b/core/src/main/java/ai/z/openapi/service/model/ToolCalls.java @@ -27,9 +27,12 @@ public class ToolCalls { private String id; /** - * Type of tool called by the model, currently only supports 'function'. + * Type of tool called by the model, currently only supports 'function', 'mcp'. */ @JsonProperty("type") private String type; + @JsonProperty("mcp") + private MCPToolCall mcp; + } diff --git a/core/src/test/java/ai/z/openapi/service/chat/ChatServiceTest.java b/core/src/test/java/ai/z/openapi/service/chat/ChatServiceTest.java index c6cc14d..ebf37a1 100644 --- a/core/src/test/java/ai/z/openapi/service/chat/ChatServiceTest.java +++ b/core/src/test/java/ai/z/openapi/service/chat/ChatServiceTest.java @@ -3,6 +3,7 @@ package ai.z.openapi.service.chat; import ai.z.openapi.ZaiClient; import ai.z.openapi.core.Constants; import ai.z.openapi.core.config.ZaiConfig; +import ai.z.openapi.service.model.*; import ai.z.openapi.service.model.AsyncResultRetrieveParams; import ai.z.openapi.service.model.ChatCompletionCreateParams; import ai.z.openapi.service.model.ChatCompletionResponse; @@ -161,6 +162,86 @@ public class ChatServiceTest { logger.info("Stream chat completion test completed"); } + @Test + @DisplayName("Test Synchronous Chat Completion - MCP Tool") + @EnabledIfEnvironmentVariable(named = "ZAI_API_KEY", matches = "^[^.]+\\.[^.]+$") + void testSyncChatCompletion_MCP_ServerUrl() throws JsonProcessingException { + // Prepare test data + List messages = new ArrayList<>(); + ChatMessage userMessage = new ChatMessage(ChatMessageRole.USER.value(), "Hello, please introduce GPT?"); + messages.add(userMessage); + + String requestId = String.format(REQUEST_ID_TEMPLATE, System.currentTimeMillis()); + + ChatCompletionCreateParams request = ChatCompletionCreateParams.builder() + .model(Constants.ModelChatGLM4) + .stream(Boolean.FALSE) + .messages(messages) + .requestId(requestId) + .tools(Arrays.asList(ChatTool.builder() + .type(ChatToolType.MCP.value()) + .mcp(MCPTool.builder() + .server_url("https://open.bigmodel.cn/api/mcp/sogou/sse") + .server_label("sougou") + .transport_type("sse") + .headers(Collections.singletonMap("Authorization", "Bearer " + System.getProperty("ZAI_API_KEY"))) + .build()) + .build())) + .build(); + + // Execute test + ChatCompletionResponse response = chatService.createChatCompletion(request); + + // Verify results + assertNotNull(response, "Response should not be null"); + assertTrue(response.isSuccess(), "Response should be successful"); + assertNotNull(response.getData(), "Response data should not be null"); + assertEquals(requestId, response.getData().getRequestId(), "Request ID should match"); + assertNotNull(response.getData().getChoices(), "Response data should not be null"); + assertFalse(response.getData().getChoices().isEmpty(), "Response data should not be empty"); + assertNull(response.getError(), "Response error should be null"); + logger.info("Synchronous chat completion response: {}", mapper.writeValueAsString(response)); + } + + @Test + @DisplayName("Test Synchronous Chat Completion - MCP Tool") + @EnabledIfEnvironmentVariable(named = "ZAI_API_KEY", matches = "^[^.]+\\.[^.]+$") + void testSyncChatCompletion_MCP_ServerLabel() throws JsonProcessingException { + // Prepare test data + List messages = new ArrayList<>(); + ChatMessage userMessage = new ChatMessage(ChatMessageRole.USER.value(), "Hello, please introduce GPT?"); + messages.add(userMessage); + + String requestId = String.format(REQUEST_ID_TEMPLATE, System.currentTimeMillis()); + + ChatCompletionCreateParams request = ChatCompletionCreateParams.builder() + .model(Constants.ModelChatGLM4) + .stream(Boolean.FALSE) + .messages(messages) + .requestId(requestId) + .tools(Arrays.asList(ChatTool.builder() + .type(ChatToolType.MCP.value()) + .mcp(MCPTool.builder() + .server_label("sougou_search") + .headers(Collections.singletonMap("Authorization", "Bearer " + System.getProperty("ZAI_API_KEY"))) + .build()) + .build())) + .build(); + + // Execute test + ChatCompletionResponse response = chatService.createChatCompletion(request); + + // Verify results + assertNotNull(response, "Response should not be null"); + assertTrue(response.isSuccess(), "Response should be successful"); + assertNotNull(response.getData(), "Response data should not be null"); + assertEquals(requestId, response.getData().getRequestId(), "Request ID should match"); + assertNotNull(response.getData().getChoices(), "Response data should not be null"); + assertFalse(response.getData().getChoices().isEmpty(), "Response data should not be empty"); + assertNull(response.getError(), "Response error should be null"); + logger.info("Synchronous chat completion response: {}", mapper.writeValueAsString(response)); + } + @Test @DisplayName("Test Asynchronous Chat Completion") @EnabledIfEnvironmentVariable(named = "ZAI_API_KEY", matches = "^[^.]+\\.[^.]+$")