Improvements to end to end test and test stdio mcp client

This commit is contained in:
Daniel Allen
2026-03-09 08:53:14 +00:00
parent 67282afa91
commit 6b7ebb6903
2 changed files with 44 additions and 61 deletions
@@ -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")
}
} }
} }