mirror of
https://github.com/PortSwigger/mcp-server
synced 2026-06-21 13:45:21 +00:00
Improvements to end to end test and test stdio mcp client
This commit is contained in:
@@ -3,9 +3,9 @@ package net.portswigger.mcp
|
|||||||
import burp.api.montoya.MontoyaApi
|
import burp.api.montoya.MontoyaApi
|
||||||
import burp.api.montoya.logging.Logging
|
import burp.api.montoya.logging.Logging
|
||||||
import burp.api.montoya.persistence.PersistedObject
|
import burp.api.montoya.persistence.PersistedObject
|
||||||
import io.modelcontextprotocol.kotlin.sdk.TextContent
|
|
||||||
import io.mockk.every
|
import io.mockk.every
|
||||||
import io.mockk.mockk
|
import io.mockk.mockk
|
||||||
|
import io.modelcontextprotocol.kotlin.sdk.TextContent
|
||||||
import kotlinx.coroutines.delay
|
import kotlinx.coroutines.delay
|
||||||
import kotlinx.coroutines.runBlocking
|
import kotlinx.coroutines.runBlocking
|
||||||
import net.portswigger.mcp.config.McpConfig
|
import net.portswigger.mcp.config.McpConfig
|
||||||
@@ -13,11 +13,13 @@ import org.junit.jupiter.api.AfterEach
|
|||||||
import org.junit.jupiter.api.Assertions.*
|
import org.junit.jupiter.api.Assertions.*
|
||||||
import org.junit.jupiter.api.BeforeEach
|
import org.junit.jupiter.api.BeforeEach
|
||||||
import org.junit.jupiter.api.Test
|
import org.junit.jupiter.api.Test
|
||||||
|
import org.junit.jupiter.api.Timeout
|
||||||
import org.junit.jupiter.api.assertDoesNotThrow
|
import org.junit.jupiter.api.assertDoesNotThrow
|
||||||
import org.slf4j.LoggerFactory
|
import org.slf4j.LoggerFactory
|
||||||
import java.io.File
|
import java.io.File
|
||||||
import java.net.ServerSocket
|
import java.net.ServerSocket
|
||||||
import kotlin.time.Duration.Companion.seconds
|
import java.util.concurrent.TimeUnit
|
||||||
|
import kotlin.time.Duration.Companion.milliseconds
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* End-to-end test verifying the full stack:
|
* End-to-end test verifying the full stack:
|
||||||
@@ -25,6 +27,7 @@ import kotlin.time.Duration.Companion.seconds
|
|||||||
*
|
*
|
||||||
* Requires libs/mcp-proxy-all.jar to be present (built via `./gradlew embedProxyJar` from root).
|
* Requires libs/mcp-proxy-all.jar to be present (built via `./gradlew embedProxyJar` from root).
|
||||||
*/
|
*/
|
||||||
|
@Timeout(30, unit = TimeUnit.SECONDS)
|
||||||
class ProxyEndToEndTest {
|
class ProxyEndToEndTest {
|
||||||
private val logger = LoggerFactory.getLogger(ProxyEndToEndTest::class.java)
|
private val logger = LoggerFactory.getLogger(ProxyEndToEndTest::class.java)
|
||||||
|
|
||||||
@@ -32,6 +35,8 @@ class ProxyEndToEndTest {
|
|||||||
private val serverManager = KtorServerManager(api)
|
private val serverManager = KtorServerManager(api)
|
||||||
private val testPort = findAvailablePort()
|
private val testPort = findAvailablePort()
|
||||||
private val persistedObject = mockk<PersistedObject>()
|
private val persistedObject = mockk<PersistedObject>()
|
||||||
|
|
||||||
|
@Volatile
|
||||||
private var serverStarted = false
|
private var serverStarted = false
|
||||||
|
|
||||||
init {
|
init {
|
||||||
@@ -76,7 +81,7 @@ class ProxyEndToEndTest {
|
|||||||
|
|
||||||
val jarFile = File("libs/mcp-proxy-all.jar")
|
val jarFile = File("libs/mcp-proxy-all.jar")
|
||||||
check(jarFile.exists()) {
|
check(jarFile.exists()) {
|
||||||
"libs/mcp-proxy-all.jar not found. Build it first: ./gradlew embedProxyJar (from repo root)"
|
"libs/mcp-proxy-all.jar not found. Build it and copy it to libs first: ./gradlew embedProxyJar (from proxy repo root)"
|
||||||
}
|
}
|
||||||
|
|
||||||
proxyProcess = ProcessBuilder(
|
proxyProcess = ProcessBuilder(
|
||||||
@@ -88,13 +93,27 @@ class ProxyEndToEndTest {
|
|||||||
"http://127.0.0.1:$testPort"
|
"http://127.0.0.1:$testPort"
|
||||||
).redirectError(ProcessBuilder.Redirect.INHERIT).start()
|
).redirectError(ProcessBuilder.Redirect.INHERIT).start()
|
||||||
|
|
||||||
delay(3.seconds)
|
|
||||||
|
|
||||||
client = TestStdioMcpClient()
|
client = TestStdioMcpClient()
|
||||||
client.connectToServer(proxyProcess.inputStream, proxyProcess.outputStream)
|
connectClientWithRetry()
|
||||||
logger.info("Test client connected to proxy on port $testPort")
|
logger.info("Test client connected to proxy on port $testPort")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
private suspend fun connectClientWithRetry() {
|
||||||
|
val maxAttempts = 10
|
||||||
|
val retryDelay = 500.milliseconds
|
||||||
|
for (attempt in 1..maxAttempts) {
|
||||||
|
check(proxyProcess.isAlive) { "Proxy process died during startup" }
|
||||||
|
try {
|
||||||
|
client.connectToServer(proxyProcess.inputStream, proxyProcess.outputStream)
|
||||||
|
return
|
||||||
|
} catch (e: Exception) {
|
||||||
|
if (attempt == maxAttempts) throw e
|
||||||
|
logger.info("Proxy not ready (attempt $attempt/$maxAttempts), retrying...")
|
||||||
|
delay(retryDelay)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
@AfterEach
|
@AfterEach
|
||||||
fun tearDown(): Unit = runBlocking {
|
fun tearDown(): Unit = runBlocking {
|
||||||
try {
|
try {
|
||||||
@@ -108,8 +127,7 @@ class ProxyEndToEndTest {
|
|||||||
try {
|
try {
|
||||||
if (::proxyProcess.isInitialized) {
|
if (::proxyProcess.isInitialized) {
|
||||||
proxyProcess.destroy()
|
proxyProcess.destroy()
|
||||||
if (proxyProcess.isAlive) {
|
if (!proxyProcess.waitFor(2, TimeUnit.SECONDS)) {
|
||||||
delay(1000)
|
|
||||||
proxyProcess.destroyForcibly()
|
proxyProcess.destroyForcibly()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -117,7 +135,7 @@ class ProxyEndToEndTest {
|
|||||||
logger.warn("Error destroying proxy process: ${e.message}")
|
logger.warn("Error destroying proxy process: ${e.message}")
|
||||||
}
|
}
|
||||||
|
|
||||||
serverManager.stop {}
|
serverManager.shutdown()
|
||||||
}
|
}
|
||||||
|
|
||||||
@Test
|
@Test
|
||||||
@@ -142,6 +160,7 @@ class ProxyEndToEndTest {
|
|||||||
val result = client.callTool("url_encode", mapOf("content" to "hello world"))
|
val result = client.callTool("url_encode", mapOf("content" to "hello world"))
|
||||||
assertNotNull(result, "Tool call result should not be null")
|
assertNotNull(result, "Tool call result should not be null")
|
||||||
assertFalse(result?.isError ?: true, "Tool call should not return an error")
|
assertFalse(result?.isError ?: true, "Tool call should not return an error")
|
||||||
|
assertTrue(result?.content?.first() is TextContent, "Result should contain TextContent")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,9 @@
|
|||||||
package net.portswigger.mcp
|
package net.portswigger.mcp
|
||||||
|
|
||||||
import io.modelcontextprotocol.kotlin.sdk.*
|
import io.modelcontextprotocol.kotlin.sdk.CallToolResultBase
|
||||||
|
import io.modelcontextprotocol.kotlin.sdk.EmptyRequestResult
|
||||||
|
import io.modelcontextprotocol.kotlin.sdk.Implementation
|
||||||
|
import io.modelcontextprotocol.kotlin.sdk.Tool
|
||||||
import io.modelcontextprotocol.kotlin.sdk.client.Client
|
import io.modelcontextprotocol.kotlin.sdk.client.Client
|
||||||
import io.modelcontextprotocol.kotlin.sdk.client.StdioClientTransport
|
import io.modelcontextprotocol.kotlin.sdk.client.StdioClientTransport
|
||||||
import kotlinx.io.asSink
|
import kotlinx.io.asSink
|
||||||
@@ -14,69 +17,30 @@ class TestStdioMcpClient {
|
|||||||
private val logger = LoggerFactory.getLogger(TestStdioMcpClient::class.java)
|
private val logger = LoggerFactory.getLogger(TestStdioMcpClient::class.java)
|
||||||
private val mcp: Client = Client(clientInfo = Implementation(name = "test-mcp-client", version = "1.0.0"))
|
private val mcp: Client = Client(clientInfo = Implementation(name = "test-mcp-client", version = "1.0.0"))
|
||||||
|
|
||||||
private lateinit var tools: List<Tool>
|
suspend fun connectToServer(input: InputStream, output: OutputStream) {
|
||||||
private lateinit var input: InputStream
|
val transport = StdioClientTransport(
|
||||||
private lateinit var output: OutputStream
|
input = input.asSource().buffered(),
|
||||||
|
output = output.asSink().buffered()
|
||||||
|
)
|
||||||
|
|
||||||
suspend fun connectToServer(input: InputStream = System.`in`, output: OutputStream = System.out) {
|
mcp.connect(transport)
|
||||||
try {
|
logger.info("Connected to server")
|
||||||
this.input = input
|
|
||||||
this.output = output
|
|
||||||
|
|
||||||
val transport = StdioClientTransport(
|
|
||||||
input = input.asSource().buffered(),
|
|
||||||
output = output.asSink().buffered()
|
|
||||||
)
|
|
||||||
|
|
||||||
mcp.connect(transport)
|
|
||||||
|
|
||||||
val toolsResult = mcp.listTools()
|
|
||||||
tools = toolsResult?.tools ?: emptyList()
|
|
||||||
println("Connected to server with tools: ${tools.joinToString(", ") { it.name }}")
|
|
||||||
} catch (e: Exception) {
|
|
||||||
println("Failed to connect to MCP server: $e")
|
|
||||||
throw e
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
suspend fun ping(): EmptyRequestResult {
|
suspend fun ping(): EmptyRequestResult {
|
||||||
try {
|
return mcp.ping()
|
||||||
val pingRequest = mcp.ping()
|
|
||||||
logger.info("Ping sent: $pingRequest")
|
|
||||||
return pingRequest
|
|
||||||
} catch (e: Exception) {
|
|
||||||
logger.error("Failed to send ping: $e")
|
|
||||||
throw e
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
suspend fun listTools(): List<Tool> {
|
suspend fun listTools(): List<Tool> {
|
||||||
try {
|
return mcp.listTools().tools
|
||||||
val toolsResult = mcp.listTools()
|
|
||||||
tools = toolsResult?.tools ?: emptyList()
|
|
||||||
logger.info("Tools listed: ${tools.joinToString(", ") { it.name }}")
|
|
||||||
return tools
|
|
||||||
} catch (e: Exception) {
|
|
||||||
logger.error("Failed to list tools: $e")
|
|
||||||
throw e
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
suspend fun callTool(toolName: String, arguments: Map<String, Any>): CallToolResultBase? {
|
suspend fun callTool(toolName: String, arguments: Map<String, Any>): CallToolResultBase? {
|
||||||
try {
|
return mcp.callTool(toolName, arguments)
|
||||||
return mcp.callTool(toolName, arguments)
|
|
||||||
} catch (e: Exception) {
|
|
||||||
logger.error("Failed to call tool: $e")
|
|
||||||
throw e
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
suspend fun close() {
|
suspend fun close() {
|
||||||
try {
|
mcp.close()
|
||||||
mcp.close()
|
logger.info("MCP client closed successfully.")
|
||||||
logger.info("MCP client closed successfully.")
|
|
||||||
} catch (e: Exception) {
|
|
||||||
logger.error("Failed to close MCP client: $e")
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user