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.logging.Logging
|
||||
import burp.api.montoya.persistence.PersistedObject
|
||||
import io.modelcontextprotocol.kotlin.sdk.TextContent
|
||||
import io.mockk.every
|
||||
import io.mockk.mockk
|
||||
import io.modelcontextprotocol.kotlin.sdk.TextContent
|
||||
import kotlinx.coroutines.delay
|
||||
import kotlinx.coroutines.runBlocking
|
||||
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.BeforeEach
|
||||
import org.junit.jupiter.api.Test
|
||||
import org.junit.jupiter.api.Timeout
|
||||
import org.junit.jupiter.api.assertDoesNotThrow
|
||||
import org.slf4j.LoggerFactory
|
||||
import java.io.File
|
||||
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:
|
||||
@@ -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).
|
||||
*/
|
||||
@Timeout(30, unit = TimeUnit.SECONDS)
|
||||
class ProxyEndToEndTest {
|
||||
private val logger = LoggerFactory.getLogger(ProxyEndToEndTest::class.java)
|
||||
|
||||
@@ -32,6 +35,8 @@ class ProxyEndToEndTest {
|
||||
private val serverManager = KtorServerManager(api)
|
||||
private val testPort = findAvailablePort()
|
||||
private val persistedObject = mockk<PersistedObject>()
|
||||
|
||||
@Volatile
|
||||
private var serverStarted = false
|
||||
|
||||
init {
|
||||
@@ -76,7 +81,7 @@ class ProxyEndToEndTest {
|
||||
|
||||
val jarFile = File("libs/mcp-proxy-all.jar")
|
||||
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(
|
||||
@@ -88,13 +93,27 @@ class ProxyEndToEndTest {
|
||||
"http://127.0.0.1:$testPort"
|
||||
).redirectError(ProcessBuilder.Redirect.INHERIT).start()
|
||||
|
||||
delay(3.seconds)
|
||||
|
||||
client = TestStdioMcpClient()
|
||||
client.connectToServer(proxyProcess.inputStream, proxyProcess.outputStream)
|
||||
connectClientWithRetry()
|
||||
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
|
||||
fun tearDown(): Unit = runBlocking {
|
||||
try {
|
||||
@@ -108,8 +127,7 @@ class ProxyEndToEndTest {
|
||||
try {
|
||||
if (::proxyProcess.isInitialized) {
|
||||
proxyProcess.destroy()
|
||||
if (proxyProcess.isAlive) {
|
||||
delay(1000)
|
||||
if (!proxyProcess.waitFor(2, TimeUnit.SECONDS)) {
|
||||
proxyProcess.destroyForcibly()
|
||||
}
|
||||
}
|
||||
@@ -117,7 +135,7 @@ class ProxyEndToEndTest {
|
||||
logger.warn("Error destroying proxy process: ${e.message}")
|
||||
}
|
||||
|
||||
serverManager.stop {}
|
||||
serverManager.shutdown()
|
||||
}
|
||||
|
||||
@Test
|
||||
@@ -142,6 +160,7 @@ class ProxyEndToEndTest {
|
||||
val result = client.callTool("url_encode", mapOf("content" to "hello world"))
|
||||
assertNotNull(result, "Tool call result should not be null")
|
||||
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
|
||||
|
||||
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.StdioClientTransport
|
||||
import kotlinx.io.asSink
|
||||
@@ -14,69 +17,30 @@ class TestStdioMcpClient {
|
||||
private val logger = LoggerFactory.getLogger(TestStdioMcpClient::class.java)
|
||||
private val mcp: Client = Client(clientInfo = Implementation(name = "test-mcp-client", version = "1.0.0"))
|
||||
|
||||
private lateinit var tools: List<Tool>
|
||||
private lateinit var input: InputStream
|
||||
private lateinit var output: OutputStream
|
||||
suspend fun connectToServer(input: InputStream, output: OutputStream) {
|
||||
val transport = StdioClientTransport(
|
||||
input = input.asSource().buffered(),
|
||||
output = output.asSink().buffered()
|
||||
)
|
||||
|
||||
suspend fun connectToServer(input: InputStream = System.`in`, output: OutputStream = System.out) {
|
||||
try {
|
||||
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
|
||||
}
|
||||
mcp.connect(transport)
|
||||
logger.info("Connected to server")
|
||||
}
|
||||
|
||||
suspend fun ping(): EmptyRequestResult {
|
||||
try {
|
||||
val pingRequest = mcp.ping()
|
||||
logger.info("Ping sent: $pingRequest")
|
||||
return pingRequest
|
||||
} catch (e: Exception) {
|
||||
logger.error("Failed to send ping: $e")
|
||||
throw e
|
||||
}
|
||||
return mcp.ping()
|
||||
}
|
||||
|
||||
suspend fun listTools(): List<Tool> {
|
||||
try {
|
||||
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
|
||||
}
|
||||
return mcp.listTools().tools
|
||||
}
|
||||
|
||||
suspend fun callTool(toolName: String, arguments: Map<String, Any>): CallToolResultBase? {
|
||||
try {
|
||||
return mcp.callTool(toolName, arguments)
|
||||
} catch (e: Exception) {
|
||||
logger.error("Failed to call tool: $e")
|
||||
throw e
|
||||
}
|
||||
return mcp.callTool(toolName, arguments)
|
||||
}
|
||||
|
||||
suspend fun close() {
|
||||
try {
|
||||
mcp.close()
|
||||
logger.info("MCP client closed successfully.")
|
||||
} catch (e: Exception) {
|
||||
logger.error("Failed to close MCP client: $e")
|
||||
}
|
||||
mcp.close()
|
||||
logger.info("MCP client closed successfully.")
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user