fix: resolve MCP Gateway SSE hanging issue by replacing RouterFunction with SseEmitter in RestController

This commit is contained in:
jade
2026-07-16 10:56:13 +09:00
parent 6d49edfc2b
commit e5f6321eed
4 changed files with 202 additions and 73 deletions

View File

@@ -0,0 +1,146 @@
package io.shinhanlife.axhub.biz.mcp.gateway.sync;
import com.fasterxml.jackson.databind.ObjectMapper;
import io.modelcontextprotocol.json.TypeRef;
import io.modelcontextprotocol.spec.McpSchema;
import io.modelcontextprotocol.spec.McpServerSession;
import io.modelcontextprotocol.spec.McpServerTransport;
import io.modelcontextprotocol.spec.McpServerTransportProvider;
import org.springframework.http.MediaType;
import org.springframework.web.servlet.function.RouterFunction;
import org.springframework.web.servlet.function.RouterFunctions;
import org.springframework.web.servlet.function.ServerResponse;
import reactor.core.publisher.Mono;
import java.util.Map;
import java.util.UUID;
import java.util.concurrent.ConcurrentHashMap;
import static org.springframework.web.servlet.function.RequestPredicates.GET;
import static org.springframework.web.servlet.function.RequestPredicates.POST;
import static org.springframework.web.servlet.function.RequestPredicates.accept;
public class CustomWebMvcSseServerTransportProvider implements McpServerTransportProvider {
private McpServerSession.Factory sessionFactory;
private final String sseEndpoint;
private final String messageEndpoint;
private final Map<String, McpServerSession> sessions = new ConcurrentHashMap<>();
private final ObjectMapper objectMapper;
public CustomWebMvcSseServerTransportProvider(String sseEndpoint, String messageEndpoint, ObjectMapper objectMapper) {
this.sseEndpoint = sseEndpoint;
this.messageEndpoint = messageEndpoint;
this.objectMapper = objectMapper != null ? objectMapper : new ObjectMapper();
}
@Override
public void setSessionFactory(McpServerSession.Factory sessionFactory) {
this.sessionFactory = sessionFactory;
}
@Override
public Mono<Void> closeGracefully() {
return Mono.fromRunnable(() -> {
sessions.values().forEach(session -> {
try {
session.closeGracefully().subscribe();
} catch (Exception ignored) {}
});
sessions.clear();
});
}
@Override
public Mono<Void> notifyClients(String method, Object params) {
return Mono.when(sessions.values().stream()
.map(session -> session.sendNotification(method, params))
.toList());
}
public org.springframework.web.servlet.mvc.method.annotation.SseEmitter handleSse() {
if (sessionFactory == null) {
throw new IllegalStateException("SessionFactory not configured");
}
org.springframework.web.servlet.mvc.method.annotation.SseEmitter emitter = new org.springframework.web.servlet.mvc.method.annotation.SseEmitter(-1L);
String sessionId = UUID.randomUUID().toString();
CustomMcpSessionTransport sessionTransport = new CustomMcpSessionTransport(emitter, sessionId);
McpServerSession session = sessionFactory.create(sessionTransport);
sessions.put(sessionId, session);
emitter.onCompletion(() -> sessions.remove(sessionId));
emitter.onTimeout(() -> sessions.remove(sessionId));
new Thread(() -> {
try {
Thread.sleep(100);
emitter.send(org.springframework.web.servlet.mvc.method.annotation.SseEmitter.event().name("endpoint").data(messageEndpoint + "?sessionId=" + sessionId));
} catch (Exception e) {
emitter.completeWithError(e);
}
}).start();
return emitter;
}
public org.springframework.http.ResponseEntity<String> handleMessage(String sessionId, String body) {
if (sessionId == null || !sessions.containsKey(sessionId)) {
return org.springframework.http.ResponseEntity.badRequest().body("Missing or invalid sessionId");
}
McpServerSession session = sessions.get(sessionId);
try {
java.util.Map<String, Object> map = objectMapper.readValue(body, new com.fasterxml.jackson.core.type.TypeReference<java.util.Map<String, Object>>() {});
io.modelcontextprotocol.spec.McpSchema.JSONRPCMessage message;
if (map.containsKey("id")) {
if (map.containsKey("method")) {
message = objectMapper.convertValue(map, io.modelcontextprotocol.spec.McpSchema.JSONRPCRequest.class);
} else {
message = objectMapper.convertValue(map, io.modelcontextprotocol.spec.McpSchema.JSONRPCResponse.class);
}
} else {
message = objectMapper.convertValue(map, io.modelcontextprotocol.spec.McpSchema.JSONRPCNotification.class);
}
session.handle(message).subscribe();
return org.springframework.http.ResponseEntity.ok().build();
} catch (Exception e) {
return org.springframework.http.ResponseEntity.status(500).body(e.getMessage());
}
}
private class CustomMcpSessionTransport implements McpServerTransport {
private final org.springframework.web.servlet.mvc.method.annotation.SseEmitter emitter;
private final String sessionId;
public CustomMcpSessionTransport(org.springframework.web.servlet.mvc.method.annotation.SseEmitter emitter, String sessionId) {
this.emitter = emitter;
this.sessionId = sessionId;
}
@Override
public Mono<Void> sendMessage(McpSchema.JSONRPCMessage message) {
return Mono.fromRunnable(() -> {
try {
String json = objectMapper.writeValueAsString(message);
emitter.send(org.springframework.web.servlet.mvc.method.annotation.SseEmitter.event().name("message").data(json));
} catch (Exception e) {
throw new RuntimeException(e);
}
});
}
@Override
public Mono<Void> closeGracefully() {
return Mono.fromRunnable(emitter::complete);
}
@Override
public <T> T unmarshalFrom(Object object, TypeRef<T> typeRef) {
return objectMapper.convertValue(object, objectMapper.constructType(typeRef.getType()));
}
}
}

View File

@@ -0,0 +1,42 @@
package io.shinhanlife.axhub.biz.mcp.gateway.sync;
import org.springframework.http.ResponseEntity;
import org.springframework.web.bind.annotation.GetMapping;
import org.springframework.web.bind.annotation.PathVariable;
import org.springframework.web.bind.annotation.PostMapping;
import org.springframework.web.bind.annotation.RequestBody;
import org.springframework.web.bind.annotation.RequestParam;
import org.springframework.web.bind.annotation.RestController;
import org.springframework.web.servlet.mvc.method.annotation.SseEmitter;
@RestController
public class DynamicMcpController {
private final DynamicMcpServerManager manager;
public DynamicMcpController(DynamicMcpServerManager manager) {
this.manager = manager;
}
@GetMapping("/mcp/sse/{category}")
public SseEmitter handleSse(@PathVariable("category") String category) {
CustomWebMvcSseServerTransportProvider transport = manager.getTransport(category);
if (transport == null) {
throw new IllegalArgumentException("Unknown category: " + category);
}
return transport.handleSse();
}
@PostMapping("/mcp/message/{category}")
public ResponseEntity<String> handleMessage(
@PathVariable("category") String category,
@RequestParam("sessionId") String sessionId,
@RequestBody String body) {
CustomWebMvcSseServerTransportProvider transport = manager.getTransport(category);
if (transport == null) {
return ResponseEntity.badRequest().body("Unknown category: " + category);
}
return transport.handleMessage(sessionId, body);
}
}

View File

@@ -1,29 +0,0 @@
package io.shinhanlife.axhub.biz.mcp.gateway.sync;
import org.springframework.context.annotation.Bean;
import org.springframework.context.annotation.Configuration;
import org.springframework.web.servlet.function.RouterFunction;
import org.springframework.web.servlet.function.ServerResponse;
/**
* @package io.shinhanlife.axhub.biz.mcp.gateway.sync
* @className DynamicMcpRouterConfig
* @description AX HUB 시스템 처리 클래스
* @author 김형식
* @create 2026.09.01
* <pre>
* ---------- 개정이력 ----------
* 수정일 수정자 수정내용
* ---------- -------- ---------------------------
* 2026.09.01 김형식 최초생성
*
* </pre>
*/
@Configuration
public class DynamicMcpRouterConfig {
@Bean
public RouterFunction<ServerResponse> dynamicMcpRouterFunction(DynamicMcpServerManager manager) {
return manager.getDynamicRouter();
}
}

View File

@@ -5,49 +5,32 @@ import io.modelcontextprotocol.server.McpSyncServer;
import io.shinhanlife.axhub.biz.mcp.gateway.dto.ToolMetadata; import io.shinhanlife.axhub.biz.mcp.gateway.dto.ToolMetadata;
import org.slf4j.Logger; import org.slf4j.Logger;
import org.slf4j.LoggerFactory; import org.slf4j.LoggerFactory;
import org.springframework.ai.mcp.server.webmvc.transport.WebMvcSseServerTransportProvider;
import org.springframework.stereotype.Component; import org.springframework.stereotype.Component;
import org.springframework.web.servlet.function.HandlerFunction; import com.fasterxml.jackson.databind.ObjectMapper;
import org.springframework.web.servlet.function.RouterFunction;
import org.springframework.web.servlet.function.ServerResponse;
import java.util.Collections; import java.util.Collections;
import java.util.List; import java.util.List;
import java.util.Map; import java.util.Map;
import java.util.Optional;
import java.util.Set; import java.util.Set;
import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.ConcurrentHashMap;
import java.util.stream.Collectors; import java.util.stream.Collectors;
/**
* @package io.shinhanlife.axhub.biz.mcp.gateway.sync
* @className DynamicMcpServerManager
* @description AX HUB 시스템 처리 클래스
* @author 김형식
* @create 2026.09.01
* <pre>
* ---------- 개정이력 ----------
* 수정일 수정자 수정내용
* ---------- -------- ---------------------------
* 2026.09.01 김형식 최초생성
*
* </pre>
*/
@Component @Component
public class DynamicMcpServerManager { public class DynamicMcpServerManager {
private static final Logger log = LoggerFactory.getLogger(DynamicMcpServerManager.class); private static final Logger log = LoggerFactory.getLogger(DynamicMcpServerManager.class);
private final Map<String, McpSyncServer> categoryServers = new ConcurrentHashMap<>(); private final Map<String, McpSyncServer> categoryServers = new ConcurrentHashMap<>();
private final Map<String, WebMvcSseServerTransportProvider> categoryTransports = new ConcurrentHashMap<>(); private final Map<String, CustomWebMvcSseServerTransportProvider> categoryTransports = new ConcurrentHashMap<>();
private final Map<String, Set<String>> managedToolNamesPerCategory = new ConcurrentHashMap<>(); private final Map<String, Set<String>> managedToolNamesPerCategory = new ConcurrentHashMap<>();
private final RegistryMcpToolSpecificationFactory specificationFactory; private final RegistryMcpToolSpecificationFactory specificationFactory;
private final ObjectMapper objectMapper;
public DynamicMcpServerManager(RegistryMcpToolSpecificationFactory specificationFactory) { public DynamicMcpServerManager(RegistryMcpToolSpecificationFactory specificationFactory, ObjectMapper objectMapper) {
this.specificationFactory = specificationFactory; this.specificationFactory = specificationFactory;
this.objectMapper = objectMapper;
// Pre-initialize basic categories so their endpoints are always open // Pre-initialize basic categories so their endpoints are always open
// even if there are 0 tools registered in Redis initially.
getOrCreateServer("common"); getOrCreateServer("common");
getOrCreateServer("hr"); getOrCreateServer("hr");
getOrCreateServer("payment"); getOrCreateServer("payment");
@@ -58,19 +41,17 @@ public class DynamicMcpServerManager {
private McpSyncServer getOrCreateServer(String categoryKey) { private McpSyncServer getOrCreateServer(String categoryKey) {
String safeCategory = (categoryKey == null || categoryKey.trim().isEmpty()) ? "common" : categoryKey.toLowerCase(); String safeCategory = (categoryKey == null || categoryKey.trim().isEmpty()) ? "common" : categoryKey.toLowerCase();
return categoryServers.computeIfAbsent(safeCategory, key -> { return categoryServers.computeIfAbsent(safeCategory, key -> {
log.info("Creating dynamic MCP Server for category: {}", key); log.info("Creating dynamic MCP Server for category: {}", key);
String ssePath = "/mcp/sse/" + key; String ssePath = "/mcp/sse/" + key;
String msgPath = "/mcp/message/" + key; String msgPath = "/mcp/message/" + key;
WebMvcSseServerTransportProvider transport = WebMvcSseServerTransportProvider.builder() CustomWebMvcSseServerTransportProvider transport = new CustomWebMvcSseServerTransportProvider(ssePath, msgPath, objectMapper);
.sseEndpoint(ssePath)
.messageEndpoint(msgPath)
.build();
McpSyncServer newServer = McpServer.sync(transport) McpSyncServer newServer = McpServer.sync(transport)
.serverInfo("AXHUB-Gateway-" + key, "1.0.0") .serverInfo("AXHUB-Gateway-" + key, "1.0.0")
.capabilities(io.modelcontextprotocol.spec.McpSchema.ServerCapabilities.builder().tools(true).build())
.build(); .build();
categoryTransports.put(key, transport); categoryTransports.put(key, transport);
@@ -114,19 +95,8 @@ public class DynamicMcpServerManager {
return Collections.unmodifiableSet(categoryServers.keySet()); return Collections.unmodifiableSet(categoryServers.keySet());
} }
public RouterFunction<ServerResponse> getDynamicRouter() { public CustomWebMvcSseServerTransportProvider getTransport(String categoryKey) {
return request -> { String safeCategory = (categoryKey == null || categoryKey.trim().isEmpty()) ? "common" : categoryKey.toLowerCase();
log.info("[DynamicMcpRouter] Incoming request path: {}, method: {}", request.path(), request.method()); return categoryTransports.get(safeCategory);
for (Map.Entry<String, WebMvcSseServerTransportProvider> entry : categoryTransports.entrySet()) {
WebMvcSseServerTransportProvider transport = entry.getValue();
Optional<HandlerFunction<ServerResponse>> handler = transport.getRouterFunction().route(request);
if (handler.isPresent()) {
log.info("[DynamicMcpRouter] Matched handler for key: {}", entry.getKey());
return handler;
}
}
log.info("[DynamicMcpRouter] No handler matched for path: {}", request.path());
return Optional.empty();
};
} }
} }