Add DNS Rebinding Protection (#19)

* Implement CORS support and DNS rebinding protection in Ktor server

* Enhance DNS rebinding protection by refining origin and referer checks, and adding user agent validation

* Add HTTP request approval feature with auto-approval targets management

* Refactor Config UI layout and enhance component styling for better usability

* Enhance Config UI with Material Design components

* Refactor build configuration and enhance UI design consistency with shared constants

* Refactor Config UI to use shared design colors for improved consistency and readability

* Suppress unused warning in ExtensionBase class

* Refactor Config UI layout for improved alignment and spacing consistency

* Increase scroll pane dimensions in Config UI for better visibility of targets list

* Enhance HTTP request approval dialog to display request content and improve layout

* Add history access approval feature with UI options for HTTP and WebSocket

* Replace enabled checkbox with toggle switch for improved UI interaction

* Refactor Config UI for improved layout and component consistency

* Add hover effect to targets list for improved user interaction

* Enhance targets list with rollover effect and keyboard support for improved accessibility

* Add validation for IPv4 and IPv6 addresses in target input

* Improve target validation to support IPv4, IPv6, and wildcard formats

* Refactor ConfigUi

* Refactor UI components to improve color management and responsiveness

* Enhance UI responsiveness and button sizing

* Update dependencies

* Add code execution warning to config editing tooling checkbox label

* Update mcp-sdk version to 0.5.0

* Refactor dialog handling to improve usability and integrate Montoya API for HTTP request management

* Refactor dialog components and enhance config editing checkbox with subtitle

* Set version to 1.1.0

* Integrate logging into McpConfig

* Listener management in McpConfig

* Add sparkle animation effect to toggle interaction in Design component

* Improve sparkle effect positioning and component dimensions
This commit is contained in:
portswigger-penguin
2025-06-16 10:52:27 +01:00
committed by GitHub
parent aa664857f3
commit 30901e8473
30 changed files with 3340 additions and 379 deletions
+1 -1
View File
@@ -1,6 +1,6 @@
<?xml version="1.0" encoding="UTF-8"?>
<project version="4">
<component name="KotlinJpsPluginSettings">
<option name="version" value="2.1.0" />
<option name="version" value="2.1.21" />
</component>
</project>
+118 -39
View File
@@ -1,65 +1,144 @@
import java.time.Instant
plugins {
kotlin("jvm") version "2.1.0"
kotlin("plugin.serialization") version "2.1.0"
id("java")
id("io.ktor.plugin") version "3.0.0"
alias(libs.plugins.kotlin.jvm)
alias(libs.plugins.kotlin.serialization)
alias(libs.plugins.ktor)
java
}
group = "net.portswigger"
version = "1.0.1"
group = providers.gradleProperty("group").get()
version = providers.gradleProperty("version").get()
description = providers.gradleProperty("description").get()
repositories {
mavenCentral()
}
dependencies {
compileOnly("net.portswigger.burp.extensions:montoya-api:2025.2")
implementation("io.ktor:ktor-server-netty:3.1.1")
implementation("io.ktor:ktor-serialization-kotlinx-json:3.1.1")
implementation("io.ktor:ktor-server-content-negotiation:3.1.1")
implementation("org.jetbrains.kotlin:kotlin-stdlib:2.1.0")
implementation("org.jetbrains.kotlinx:kotlinx-serialization-json:1.8.0")
compileOnly(libs.burp.montoya.api)
implementation("io.modelcontextprotocol:kotlin-sdk:0.4.0")
implementation("io.ktor:ktor-server-core:3.1.1")
implementation("io.ktor:ktor-server-sse:3.1.1")
implementation(libs.bundles.ktor.server)
implementation(libs.kotlin.stdlib)
implementation(libs.kotlinx.serialization.json)
implementation(libs.mcp.kotlin.sdk)
testImplementation(kotlin("test"))
testImplementation("io.mockk:mockk:1.13.17")
testImplementation("net.portswigger.burp.extensions:montoya-api:2025.2")
testImplementation(libs.bundles.test.framework)
testImplementation(libs.bundles.ktor.test)
testImplementation(libs.burp.montoya.api)
}
tasks.test {
useJUnitPlatform()
java {
toolchain {
languageVersion.set(JavaLanguageVersion.of(providers.gradleProperty("java.toolchain.version").get().toInt()))
}
}
kotlin {
jvmToolchain(21)
jvmToolchain {
languageVersion.set(JavaLanguageVersion.of(providers.gradleProperty("java.toolchain.version").get().toInt()))
}
compilerOptions {
apiVersion.set(org.jetbrains.kotlin.gradle.dsl.KotlinVersion.KOTLIN_2_1)
languageVersion.set(org.jetbrains.kotlin.gradle.dsl.KotlinVersion.KOTLIN_2_1)
jvmTarget.set(org.jetbrains.kotlin.gradle.dsl.JvmTarget.JVM_21)
freeCompilerArgs.addAll(
"-Xjsr305=strict"
)
}
}
application {
mainClass.set("net.portswigger.mcp.ExtensionBase")
}
tasks.jar {
manifest {
attributes["Class-Path"] = configurations.runtimeClasspath.get().files.joinToString(" ") { it.name }
}
from({
configurations.runtimeClasspath.get().map { if (it.isDirectory) it else zipTree(it) }
})
tasks {
test {
useJUnitPlatform()
systemProperty("file.encoding", "UTF-8")
duplicatesStrategy = DuplicatesStrategy.EXCLUDE
}
tasks.register("embedProxyJar") {
dependsOn("shadowJar")
doLast {
val shadowJarFile = tasks.shadowJar.get().archiveFile.get().asFile
exec {
workingDir(projectDir)
commandLine("jar", "uf", shadowJarFile.absolutePath, "-C", "libs", "mcp-proxy-all.jar")
testLogging {
events("passed", "skipped", "failed")
showExceptions = true
showCauses = true
showStackTraces = true
}
}
jar {
enabled = false
}
shadowJar {
archiveClassifier.set("")
mergeServiceFiles()
manifest {
attributes(
mapOf(
"Implementation-Title" to project.name,
"Implementation-Version" to project.version,
"Implementation-Vendor" to "PortSwigger",
"Built-By" to System.getProperty("user.name"),
"Built-Date" to Instant.now().toString(),
"Built-JDK" to "${System.getProperty("java.version")} (${System.getProperty("java.vendor")} ${
System.getProperty("java.vm.version")
})",
"Created-By" to "Gradle ${gradle.gradleVersion}"
)
)
}
exclude("META-INF/*.SF")
exclude("META-INF/*.DSA")
exclude("META-INF/*.RSA")
exclude("META-INF/INDEX.LIST")
exclude("META-INF/DEPENDENCIES")
exclude("META-INF/NOTICE*")
exclude("META-INF/LICENSE*")
exclude("module-info.class")
duplicatesStrategy = DuplicatesStrategy.EXCLUDE
}
register("embedProxyJar") {
group = "build"
description = "Embeds the MCP proxy JAR into the shadow JAR"
dependsOn(shadowJar)
notCompatibleWithConfigurationCache("Task references other tasks at execution time")
doLast {
val shadowJarFile = shadowJar.get().archiveFile.get().asFile
val libsDir = layout.projectDirectory.dir("libs").asFile
val proxyJarFile = File(libsDir, "mcp-proxy-all.jar")
if (!proxyJarFile.exists()) {
throw GradleException("Proxy JAR not found at: ${proxyJarFile.absolutePath}")
}
exec {
workingDir(layout.projectDirectory.asFile)
commandLine("jar", "uf", shadowJarFile.absolutePath, "-C", libsDir.absolutePath, proxyJarFile.name)
}
logger.lifecycle("Embedded proxy JAR into ${shadowJarFile.name}")
}
}
build {
dependsOn(shadowJar)
}
withType<AbstractArchiveTask>().configureEach {
isPreserveFileTimestamps = false
isReproducibleFileOrder = true
}
}
tasks.wrapper {
gradleVersion = "8.10"
distributionType = Wrapper.DistributionType.BIN
}
+8
View File
@@ -1 +1,9 @@
kotlin.code.style=official
kotlin.stdlib.default.dependency=false
group=net.portswigger
version=1.1.0
description=Burp MCP Server Extension
org.gradle.parallel=true
org.gradle.caching=true
org.gradle.configuration-cache=true
java.toolchain.version=21
+62
View File
@@ -0,0 +1,62 @@
[versions]
# Build System
kotlin = "2.1.21"
ktor = "3.1.3"
# Runtime Dependencies
kotlinx-serialization = "1.8.1"
mcp-sdk = "0.5.0"
burp-montoya = "2025.5"
# Test Dependencies
mockk = "1.14.2"
[libraries]
# Kotlin
kotlin-stdlib = { module = "org.jetbrains.kotlin:kotlin-stdlib", version.ref = "kotlin" }
kotlinx-serialization-json = { module = "org.jetbrains.kotlinx:kotlinx-serialization-json", version.ref = "kotlinx-serialization" }
# Ktor Server
ktor-server-core = { module = "io.ktor:ktor-server-core", version.ref = "ktor" }
ktor-server-netty = { module = "io.ktor:ktor-server-netty", version.ref = "ktor" }
ktor-server-content-negotiation = { module = "io.ktor:ktor-server-content-negotiation", version.ref = "ktor" }
ktor-server-cors = { module = "io.ktor:ktor-server-cors", version.ref = "ktor" }
ktor-server-sse = { module = "io.ktor:ktor-server-sse", version.ref = "ktor" }
ktor-serialization-kotlinx-json = { module = "io.ktor:ktor-serialization-kotlinx-json", version.ref = "ktor" }
# MCP
mcp-kotlin-sdk = { module = "io.modelcontextprotocol:kotlin-sdk", version.ref = "mcp-sdk" }
# Montoya
burp-montoya-api = { module = "net.portswigger.burp.extensions:montoya-api", version.ref = "burp-montoya" }
# Test Dependencies
kotlin-test = { module = "org.jetbrains.kotlin:kotlin-test", version.ref = "kotlin" }
mockk = { module = "io.mockk:mockk", version.ref = "mockk" }
ktor-server-test-host = { module = "io.ktor:ktor-server-test-host", version.ref = "ktor" }
ktor-client-content-negotiation = { module = "io.ktor:ktor-client-content-negotiation", version.ref = "ktor" }
[bundles]
ktor-server = [
"ktor-server-core",
"ktor-server-netty",
"ktor-server-content-negotiation",
"ktor-server-cors",
"ktor-server-sse",
"ktor-serialization-kotlinx-json"
]
ktor-test = [
"ktor-server-test-host",
"ktor-client-content-negotiation"
]
test-framework = [
"kotlin-test",
"mockk"
]
[plugins]
kotlin-jvm = { id = "org.jetbrains.kotlin.jvm", version.ref = "kotlin" }
kotlin-serialization = { id = "org.jetbrains.kotlin.plugin.serialization", version.ref = "kotlin" }
ktor = { id = "io.ktor.plugin", version.ref = "ktor" }
Binary file not shown.
+9 -1
View File
@@ -1,4 +1,12 @@
pluginManagement {
repositories {
gradlePluginPortal()
mavenCentral()
}
}
plugins {
id("org.gradle.toolchains.foojay-resolver-convention") version "0.8.0"
}
rootProject.name = "burp-mcp"
rootProject.name = "burp-mcp"
@@ -9,30 +9,26 @@ import net.portswigger.mcp.providers.ManualProxyInstallerProvider
import net.portswigger.mcp.providers.ProxyJarManager
import net.portswigger.mcp.server.KtorServerManager
@Suppress("unused")
class ExtensionBase : BurpExtension {
override fun initialize(api: MontoyaApi) {
api.extension().setName("Burp MCP Server")
val config = McpConfig(api.persistence().extensionData())
val config = McpConfig(api.persistence().extensionData(), api.logging())
val serverManager = KtorServerManager(api)
val proxyJarManager = ProxyJarManager(api.logging())
val configUi = ConfigUi(
config = config,
providers = listOf(
config = config, providers = listOf(
ClaudeDesktopProvider(api.logging(), proxyJarManager),
ManualProxyInstallerProvider(api.logging(), proxyJarManager),
)
)
configUi.onEnabledToggled { enabled ->
val currentConfig = configUi.getConfig()
config.enabled = enabled
config.host = currentConfig.host
config.port = currentConfig.port
configUi.getConfig()
if (enabled) {
serverManager.start(config) { state ->
@@ -49,6 +45,8 @@ class ExtensionBase : BurpExtension {
api.extension().registerUnloadingHandler {
serverManager.shutdown()
configUi.cleanup()
config.cleanup()
}
if (config.enabled) {
@@ -1,8 +1,13 @@
package net.portswigger.mcp.server
import burp.api.montoya.MontoyaApi
import io.ktor.http.*
import io.ktor.server.application.*
import io.ktor.server.engine.*
import io.ktor.server.netty.*
import io.ktor.server.plugins.cors.routing.*
import io.ktor.server.request.*
import io.ktor.server.response.*
import io.modelcontextprotocol.kotlin.sdk.Implementation
import io.modelcontextprotocol.kotlin.sdk.ServerCapabilities
import io.modelcontextprotocol.kotlin.sdk.server.Server
@@ -30,8 +35,7 @@ class KtorServerManager(private val api: MontoyaApi) : ServerManager {
server = null
val mcpServer = Server(
serverInfo = Implementation("burp-suite", "1.0.0"),
options = ServerOptions(
serverInfo = Implementation("burp-suite", "1.0.0"), options = ServerOptions(
capabilities = ServerCapabilities(
tools = ServerCapabilities.Tools(listChanged = false)
)
@@ -39,6 +43,60 @@ class KtorServerManager(private val api: MontoyaApi) : ServerManager {
)
server = embeddedServer(Netty, port = config.port, host = config.host) {
install(CORS) {
allowHost("localhost:${config.port}")
allowHost("127.0.0.1:${config.port}")
allowMethod(HttpMethod.Get)
allowMethod(HttpMethod.Post)
allowMethod(HttpMethod.Options)
allowHeader(HttpHeaders.ContentType)
allowHeader(HttpHeaders.Accept)
allowHeader(HttpHeaders.CacheControl)
allowHeader("Last-Event-ID")
allowCredentials = false
allowNonSimpleContentTypes = true
maxAgeInSeconds = 3600
}
intercept(ApplicationCallPipeline.Call) {
val origin = call.request.header("Origin")
val host = call.request.header("Host")
val referer = call.request.header("Referer")
val userAgent = call.request.header("User-Agent")
if (origin != null) {
if (!isValidOrigin(origin)) {
api.logging().logToOutput("Blocked DNS rebinding attack from origin: $origin")
call.respond(HttpStatusCode.Forbidden)
return@intercept
}
} else if (isBrowserRequest(userAgent)) {
api.logging().logToOutput("Blocked browser request without Origin header")
call.respond(HttpStatusCode.Forbidden)
return@intercept
}
if (host != null && !isValidHost(host, config.port)) {
api.logging().logToOutput("Blocked DNS rebinding attack from host: $host")
call.respond(HttpStatusCode.Forbidden)
return@intercept
}
if (referer != null && !isValidReferer(referer)) {
api.logging().logToOutput("Blocked suspicious request from referer: $referer")
call.respond(HttpStatusCode.Forbidden)
return@intercept
}
call.response.header("X-Frame-Options", "DENY")
call.response.header("X-Content-Type-Options", "nosniff")
call.response.header("Referrer-Policy", "same-origin")
call.response.header("Content-Security-Policy", "default-src 'none'")
}
mcp {
mcpServer
}
@@ -81,4 +139,62 @@ class KtorServerManager(private val api: MontoyaApi) : ServerManager {
executor.shutdown()
executor.awaitTermination(10, TimeUnit.SECONDS)
}
private fun isValidOrigin(origin: String): Boolean {
try {
val url = java.net.URI(origin).toURL()
val hostname = url.host.lowercase()
val allowedHosts = setOf("localhost", "127.0.0.1")
return hostname in allowedHosts
} catch (_: Exception) {
return false
}
}
private fun isBrowserRequest(userAgent: String?): Boolean {
if (userAgent == null) return false
val userAgentLower = userAgent.lowercase()
val browserIndicators = listOf(
"mozilla/", "chrome/", "safari/", "webkit/", "gecko/", "firefox/", "edge/", "opera/", "browser"
)
return browserIndicators.any { userAgentLower.contains(it) }
}
private fun isValidHost(host: String, expectedPort: Int): Boolean {
try {
val parts = host.split(":")
val hostname = parts[0].lowercase()
val port = if (parts.size > 1) parts[1].toIntOrNull() else null
val allowedHosts = setOf("localhost", "127.0.0.1")
if (hostname !in allowedHosts) {
return false
}
if (port != null && port != expectedPort) {
return false
}
return true
} catch (_: Exception) {
return false
}
}
private fun isValidReferer(referer: String): Boolean {
try {
val url = java.net.URI(referer).toURL()
val hostname = url.host.lowercase()
val allowedHosts = setOf("localhost", "127.0.0.1")
return hostname in allowedHosts
} catch (_: Exception) {
return false
}
}
}
@@ -6,82 +6,97 @@ import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.launch
import net.portswigger.mcp.ServerState
import net.portswigger.mcp.Swing
import net.portswigger.mcp.config.components.*
import net.portswigger.mcp.providers.Provider
import java.awt.*
import java.awt.BorderLayout
import java.awt.Component.CENTER_ALIGNMENT
import java.awt.event.ItemEvent
import java.awt.GridBagLayout
import javax.swing.*
import javax.swing.BorderFactory.createEmptyBorder
import javax.swing.Box.*
import javax.swing.JOptionPane.*
import javax.swing.event.DocumentEvent
import javax.swing.event.DocumentListener
import kotlin.concurrent.thread
import javax.swing.JOptionPane.ERROR_MESSAGE
class ConfigUi(private val config: McpConfig, private val providers: List<Provider>) {
class WarningLabel(content: String = "") : JLabel(content) {
init {
foreground = UIManager.getColor("Burp.warningBarBackground")
isVisible = false
alignmentX = Component.LEFT_ALIGNMENT
}
override fun updateUI() {
super.updateUI()
foreground = UIManager.getColor("Burp.warningBarBackground")
}
}
private val panel = JPanel(BorderLayout())
val component: JComponent get() = panel
private val enabledCheckBox = JCheckBox("Enabled").apply { alignmentX = Component.LEFT_ALIGNMENT }
private val validationErrorLabel = WarningLabel()
private val listenerHandles = mutableListOf<ListenerHandle>()
private val enabledToggle: ToggleSwitch = Design.createToggleSwitch(false) { enabled ->
if (suppressToggleEvents) return@createToggleSwitch
if (enabled) {
ConfigValidation.validateServerConfig(hostField.text, portField.text)?.let { error ->
validationErrorLabel.text = error
validationErrorLabel.isVisible = true
suppressToggleEvents = true
enabledToggle.setState(false, animate = true)
suppressToggleEvents = false
return@createToggleSwitch
}
}
validationErrorLabel.isVisible = false
config.enabled = enabled
toggleListener?.invoke(enabled)
}
private val validationErrorLabel = WarningLabel()
private val hostField = JTextField(15)
private val portField = JTextField(5)
private val reinstallNotice = WarningLabel("Make sure to reinstall after changing server settings")
private lateinit var serverConfigurationPanel: ServerConfigurationPanel
private lateinit var advancedOptionsPanel: AdvancedOptionsPanel
private lateinit var autoApproveTargetsPanel: AutoApproveTargetsPanel
private lateinit var installationPanel: InstallationPanel
private var toggleListener: ((Boolean) -> Unit)? = null
private var suppressToggleEvents: Boolean = false
private var installationAvailable: Boolean = false
init {
enabledCheckBox.isSelected = config.enabled
enabledToggle.setState(config.enabled, animate = false)
hostField.text = config.host
portField.text = config.port.toString()
initializeComponents()
buildUi()
}
enabledCheckBox.addItemListener {
if (suppressToggleEvents) {
return@addItemListener
private fun initializeComponents() {
serverConfigurationPanel = ServerConfigurationPanel(
config = config, enabledToggle = enabledToggle, validationErrorLabel = validationErrorLabel
)
advancedOptionsPanel = AdvancedOptionsPanel(
hostField = hostField, portField = portField, reinstallNotice = reinstallNotice
)
autoApproveTargetsPanel = AutoApproveTargetsPanel(config = config)
installationPanel = InstallationPanel(
config = config, providers = providers, reinstallNotice = reinstallNotice, parentComponent = panel
)
setupConfigListeners()
}
private fun setupConfigListeners() {
val historyAccessRefreshListener = {
SwingUtilities.invokeLater {
serverConfigurationPanel.updateHistoryAccessCheckboxes()
}
val checked = it.stateChange == ItemEvent.SELECTED
if (checked) {
val error = getValidationError()
if (error != null) {
validationErrorLabel.text = error
validationErrorLabel.isVisible = true
suppressToggleEvents = true
enabledCheckBox.isSelected = false
suppressToggleEvents = false
return@addItemListener
}
}
validationErrorLabel.isVisible = false
toggleListener?.invoke(checked)
}
val handle = config.addHistoryAccessChangeListener(historyAccessRefreshListener)
listenerHandles.add(handle)
}
trackChanges(hostField)
trackChanges(portField)
fun cleanup() {
listenerHandles.forEach { it.remove() }
listenerHandles.clear()
if (::autoApproveTargetsPanel.isInitialized) {
autoApproveTargetsPanel.cleanup()
}
}
fun onEnabledToggled(listener: (Boolean) -> Unit) {
@@ -99,43 +114,36 @@ class ConfigUi(private val config: McpConfig, private val providers: List<Provid
suppressToggleEvents = true
val enableAdvancedOptions = state is ServerState.Stopped || state is ServerState.Failed
hostField.isEnabled = enableAdvancedOptions
portField.isEnabled = enableAdvancedOptions
installationAvailable = false
if (::advancedOptionsPanel.isInitialized) {
advancedOptionsPanel.setFieldsEnabled(enableAdvancedOptions)
}
when (state) {
ServerState.Starting, ServerState.Stopping -> {
enabledCheckBox.isEnabled = false
enabledToggle.isEnabled = false
}
ServerState.Running -> {
enabledCheckBox.isEnabled = true
enabledCheckBox.isSelected = true
installationAvailable = true
enabledToggle.isEnabled = true
enabledToggle.setState(true, animate = false)
}
ServerState.Stopped -> {
enabledCheckBox.isEnabled = true
enabledCheckBox.isSelected = false
enabledToggle.isEnabled = true
enabledToggle.setState(false, animate = false)
}
is ServerState.Failed -> {
enabledCheckBox.isEnabled = true
enabledCheckBox.isSelected = false
enabledToggle.isEnabled = true
enabledToggle.setState(false, animate = false)
val friendlyMessage = when (state.exception) {
is UnresolvedAddressException -> "Unable to resolve address"
else -> state.exception.message ?: state.exception.javaClass.simpleName
}
showMessageDialog(
panel,
"Failed to start Burp MCP Server: $friendlyMessage",
"Error",
ERROR_MESSAGE
Dialogs.showMessageDialog(
panel, "Failed to start Burp MCP Server: $friendlyMessage", ERROR_MESSAGE
)
}
}
@@ -144,214 +152,62 @@ class ConfigUi(private val config: McpConfig, private val providers: List<Provid
}
}
private fun getValidationError(): String? {
val host = hostField.text.trim()
val port = portField.text.trim().toIntOrNull()
if (host.isBlank() || !host.matches(Regex("^[a-zA-Z0-9.-]+$"))) {
return "Host must be a non-empty alphanumeric string"
}
if (port == null) {
return "Port must be a valid number"
}
if (port < 1024 || port > 65535) {
return "Port is not within valid range"
}
return null
}
private fun trackChanges(field: JTextField) {
field.document.addDocumentListener(object : DocumentListener {
override fun insertUpdate(e: DocumentEvent?) = handle()
override fun removeUpdate(e: DocumentEvent?) = handle()
override fun changedUpdate(e: DocumentEvent?) = handle()
fun handle() {
reinstallNotice.isVisible = true
}
})
}
private fun buildUi() {
val leftPanel = JPanel(GridBagLayout())
val headerBox = createVerticalBox().apply {
add(object : JLabel("Burp MCP Server") {
init {
font = font.deriveFont(Font.BOLD, 24f)
alignmentX = CENTER_ALIGNMENT
}
override fun getForeground() = UIManager.getColor("Burp.burpTitle")
})
add(createVerticalStrut(25))
add(JLabel("Burp MCP Server exposes Burp tooling to AI clients.").apply {
add(JLabel("Burp MCP Server").apply {
font = Design.Typography.headlineMedium
foreground = Design.Colors.onSurface
alignmentX = CENTER_ALIGNMENT
})
add(createVerticalStrut(15))
add(createVerticalStrut(Design.Spacing.MD))
add(JLabel("Burp MCP Server exposes Burp tooling to AI clients.").apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurfaceVariant
alignmentX = CENTER_ALIGNMENT
})
add(createVerticalStrut(Design.Spacing.MD))
add(
Anchor(
text = "Learn more about the Model Context Protocol",
url = "https://modelcontextprotocol.io/introduction"
).apply { alignmentX = CENTER_ALIGNMENT }
)
).apply { alignmentX = CENTER_ALIGNMENT })
}
leftPanel.add(headerBox)
val rightPanel = object : JPanel() {
init {
applyStyles()
}
override fun updateUI() {
super.updateUI()
applyStyles()
}
private fun applyStyles() {
background = UIManager.getColor("Burp.backgrounder")
}
}.apply {
val rightPanelContent = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
border = createEmptyBorder(15, 15, 15, 15)
}
val configEditingToolingCheckBox = JCheckBox("Enable tools that can edit your config").apply {
alignmentX = Component.LEFT_ALIGNMENT
isSelected = config.configEditingTooling
addItemListener { event -> config.configEditingTooling = event.stateChange == ItemEvent.SELECTED }
}
rightPanel.add(enabledCheckBox)
rightPanel.add(createVerticalStrut(10))
rightPanel.add(configEditingToolingCheckBox)
rightPanel.add(validationErrorLabel)
rightPanel.add(createVerticalStrut(15))
val advancedPanel = JPanel(GridBagLayout()).apply {
border = BorderFactory.createTitledBorder("Advanced options")
isOpaque = false
}
val gbc = GridBagConstraints().apply {
insets = Insets(5, 5, 5, 5)
anchor = GridBagConstraints.WEST
}
advancedPanel.add(JLabel("Server host:"), gbc)
gbc.gridx = 1
advancedPanel.add(hostField, gbc)
gbc.gridx = 0
gbc.gridy = 1
advancedPanel.add(JLabel("Server port:"), gbc)
gbc.gridx = 1
advancedPanel.add(portField, gbc)
val advancedWrapper = JPanel(FlowLayout(FlowLayout.LEFT, 0, 0)).apply {
isOpaque = false
add(advancedPanel)
alignmentX = Component.LEFT_ALIGNMENT
}
rightPanel.add(advancedWrapper)
rightPanel.add(createVerticalGlue())
rightPanel.add(reinstallNotice)
rightPanel.add(createVerticalStrut(10))
val installOptions = JPanel().apply {
layout = BoxLayout(this, BoxLayout.X_AXIS)
alignmentX = Component.LEFT_ALIGNMENT
isOpaque = false
}
providers.forEach { provider ->
val item = JButton(provider.installButtonText)
item.addActionListener {
if (!installationAvailable) {
showMessageDialog(
panel,
"Please start the Burp MCP Server first.",
"Burp MCP Server",
INFORMATION_MESSAGE
)
return@addActionListener
}
val confirmationText = provider.confirmationText
if (confirmationText != null) {
val result = showConfirmDialog(
panel,
confirmationText,
"Burp MCP Server",
YES_NO_OPTION
)
if (result != YES_OPTION) {
return@addActionListener
}
}
thread {
try {
val result = provider.install(config)
CoroutineScope(Dispatchers.Swing).launch {
reinstallNotice.isVisible = false
if (result != null) {
showMessageDialog(
panel,
result,
"Burp MCP Server",
INFORMATION_MESSAGE
)
}
}
} catch (e: Exception) {
CoroutineScope(Dispatchers.Swing).launch {
showMessageDialog(
panel,
"Failed to install for ${provider.name}: ${e.message ?: e.javaClass.simpleName}",
"${provider.name} install",
ERROR_MESSAGE
)
}
}
}
}
installOptions.add(item)
installOptions.add(createHorizontalStrut(10))
}
installOptions.add(
Anchor(
text = "Manual install steps",
url = "https://github.com/PortSwigger/mcp-server?tab=readme-ov-file#manual-installations"
background = Design.Colors.surface
border = BorderFactory.createEmptyBorder(
Design.Spacing.LG, Design.Spacing.LG, Design.Spacing.LG, Design.Spacing.LG
)
)
installOptions.maximumSize = installOptions.preferredSize
rightPanel.add(installOptions)
val columnsPanel = JPanel(GridBagLayout())
val c = GridBagConstraints().apply {
fill = GridBagConstraints.BOTH
weighty = 1.0
}
c.gridx = 0
c.gridy = 0
c.weightx = 0.35
columnsPanel.add(leftPanel, c)
val rightPanel = JScrollPane(rightPanelContent).apply {
border = null
background = Design.Colors.surface
viewport.background = Design.Colors.surface
verticalScrollBarPolicy = JScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JScrollPane.HORIZONTAL_SCROLLBAR_NEVER
verticalScrollBar.unitIncrement = 16
}
c.gridx = 1
c.weightx = 0.65
columnsPanel.add(rightPanel, c)
rightPanelContent.add(serverConfigurationPanel)
rightPanelContent.add(createVerticalStrut(Design.Spacing.LG))
rightPanelContent.add(autoApproveTargetsPanel)
rightPanelContent.add(createVerticalStrut(15))
rightPanelContent.add(advancedOptionsPanel)
rightPanelContent.add(createVerticalGlue())
rightPanelContent.add(reinstallNotice)
rightPanelContent.add(createVerticalStrut(10))
rightPanelContent.add(installationPanel)
val columnsPanel = ResponsiveColumnsPanel(leftPanel, rightPanel)
panel.add(columnsPanel, BorderLayout.CENTER)
}
}
@@ -0,0 +1,23 @@
package net.portswigger.mcp.config
object ConfigValidation {
fun validateServerConfig(host: String, portText: String): String? {
val trimmedHost = host.trim()
val port = portText.trim().toIntOrNull()
if (trimmedHost.isBlank() || !trimmedHost.matches(Regex("^[a-zA-Z0-9.-]+$"))) {
return "Host must be a non-empty alphanumeric string"
}
if (port == null) {
return "Port must be a valid number"
}
if (port < 1024 || port > 65535) {
return "Port is not within valid range"
}
return null
}
}
@@ -0,0 +1,443 @@
package net.portswigger.mcp.config
import java.awt.*
import java.awt.event.MouseAdapter
import java.awt.event.MouseEvent
import java.awt.geom.RoundRectangle2D
import javax.swing.*
import kotlin.math.cos
import kotlin.math.sin
/**
* Shared Design constants and utilities for consistent theming across the application
*/
object Design {
object Colors {
val primary: Color get() = UIManager.getColor("Burp.primaryButtonBackground") ?: Color(0xD86633)
val onPrimary: Color get() = UIManager.getColor("Burp.primaryButtonForeground") ?: Color.WHITE
val surface: Color get() = UIManager.getColor("Panel.background") ?: Color(0xFFFBFF)
val onSurface: Color get() = UIManager.getColor("Label.foreground") ?: Color(0x1A1A1A)
val onSurfaceVariant: Color get() = UIManager.getColor("Label.disabledForeground") ?: Color(0x666666)
val outline: Color get() = UIManager.getColor("Component.borderColor") ?: Color(0xCCCCCC)
val outlineVariant: Color get() = UIManager.getColor("Separator.foreground") ?: Color(0xE0E0E0)
val error: Color get() = UIManager.getColor("Burp.errorColor") ?: Color(0xB3261E)
val warning: Color get() = UIManager.getColor("Burp.warningColor") ?: Color(0xF57C00)
val transparent = Color(0, 0, 0, 0)
val listBackground: Color get() = UIManager.getColor("List.background") ?: Color.WHITE
val listSelectionBackground: Color get() = UIManager.getColor("List.selectionBackground") ?: Color(0xE3F2FD)
val listSelectionForeground: Color get() = UIManager.getColor("List.selectionForeground") ?: Color(0x1976D2)
val listHoverBackground: Color get() = UIManager.getColor("List.hoverBackground") ?: Color(0xF0F8FF)
val listAlternatingBackground: Color get() = UIManager.getColor("List.alternateRowColor") ?: Color(0xFAFAFA)
val listBorder: Color get() = UIManager.getColor("List.border") ?: Color(0xDDDDDD)
}
object Typography {
private val baseFont: Font get() = UIManager.getFont("Label.font") ?: Font("Inter", Font.PLAIN, 14)
private val baseSize: Int get() = baseFont.size
val headlineMedium: Font get() = baseFont.deriveFont(Font.BOLD, (baseSize * 2.0f))
val titleMedium: Font get() = baseFont.deriveFont(Font.BOLD, (baseSize * 1.14f))
val bodyLarge: Font get() = baseFont.deriveFont(Font.PLAIN, (baseSize * 1.14f))
val bodyMedium: Font get() = baseFont.deriveFont(Font.PLAIN, baseSize.toFloat())
val labelLarge: Font get() = baseFont.deriveFont(Font.BOLD, baseSize.toFloat())
val labelMedium: Font get() = baseFont.deriveFont(Font.BOLD, (baseSize * 0.86f))
}
object Spacing {
private val baseSize: Int get() = (UIManager.getFont("Label.font")?.size ?: 14)
private val scaleFactor: Float get() = baseSize / 14f
val SM: Int get() = (8 * scaleFactor).toInt().coerceAtLeast(4)
val MD: Int get() = (16 * scaleFactor).toInt().coerceAtLeast(8)
val LG: Int get() = (24 * scaleFactor).toInt().coerceAtLeast(12)
val XL: Int get() = (32 * scaleFactor).toInt().coerceAtLeast(16)
}
private fun calculateTextFitSize(button: JButton): Dimension {
val font = Typography.labelLarge
val metrics = button.getFontMetrics(font)
val textWidth = metrics.stringWidth(button.text)
val textHeight = metrics.height
val horizontalPadding = Spacing.LG * 2
val verticalPadding = Spacing.SM * 2 + 4
val minWidth = textWidth + horizontalPadding
val minHeight = textHeight + verticalPadding
return Dimension(
minWidth.coerceAtLeast(80),
minHeight.coerceAtLeast(40)
)
}
private fun applyButtonBaseStyle(button: JButton, customSize: Dimension?) {
button.apply {
font = Typography.labelLarge
isFocusPainted = false
cursor = Cursor.getPredefinedCursor(Cursor.HAND_CURSOR)
val textFitSize = calculateTextFitSize(this)
minimumSize = textFitSize
preferredSize = customSize ?: textFitSize
}
}
fun createFilledButton(text: String, customSize: Dimension? = null): JButton {
return object : JButton(text) {
init {
updateColorsAndSizing()
applyButtonBaseStyle(this, customSize)
}
override fun updateUI() {
super.updateUI()
updateColorsAndSizing()
applyButtonBaseStyle(this, customSize)
}
private fun updateColorsAndSizing() {
background = Colors.primary
foreground = Colors.onPrimary
border = BorderFactory.createEmptyBorder(Spacing.SM + 2, Spacing.LG, Spacing.SM + 2, Spacing.LG)
}
}
}
fun createOutlinedButton(text: String, customSize: Dimension? = null): JButton {
return object : JButton(text) {
init {
updateColorsAndSizing()
applyButtonBaseStyle(this, customSize)
}
override fun updateUI() {
super.updateUI()
updateColorsAndSizing()
applyButtonBaseStyle(this, customSize)
}
private fun updateColorsAndSizing() {
background = Colors.surface
foreground = Colors.primary
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Colors.outline, 1),
BorderFactory.createEmptyBorder(Spacing.SM + 1, Spacing.LG - 1, Spacing.SM + 1, Spacing.LG - 1)
)
}
}
}
fun createTextButton(text: String, customSize: Dimension? = null): JButton {
return object : JButton(text) {
init {
updateColorsAndSizing()
isContentAreaFilled = false
applyButtonBaseStyle(this, customSize)
}
override fun updateUI() {
super.updateUI()
updateColorsAndSizing()
applyButtonBaseStyle(this, customSize)
}
private fun updateColorsAndSizing() {
background = Colors.transparent
foreground = Colors.primary
border = BorderFactory.createEmptyBorder(Spacing.SM + 2, Spacing.LG, Spacing.SM + 2, Spacing.LG)
}
}
}
fun createToggleSwitch(initialState: Boolean = false, onToggle: (Boolean) -> Unit): ToggleSwitch {
return ToggleSwitch(initialState, onToggle)
}
fun createSectionLabel(text: String): JLabel {
return object : JLabel(text) {
init {
updateColors()
alignmentX = LEFT_ALIGNMENT
}
override fun updateUI() {
super.updateUI()
updateColors()
}
private fun updateColors() {
font = Typography.titleMedium
foreground = Colors.onSurface
}
}
}
}
class ToggleSwitch(private var isOn: Boolean, private val onToggle: (Boolean) -> Unit) : JComponent() {
companion object {
private const val TRACK_WIDTH = 44
private const val TRACK_HEIGHT = 24
private const val THUMB_SIZE = 20
private const val PADDING = 2
private const val ANIMATION_DURATION = 150
private const val TIMER_DELAY = 16
private const val SPARKLE_DURATION = 800
private const val SPARKLE_COUNT = 8
private const val SPARKLE_MARGIN = 8
private const val COMPONENT_WIDTH = TRACK_WIDTH + (SPARKLE_MARGIN * 2)
private const val COMPONENT_HEIGHT = TRACK_HEIGHT + (SPARKLE_MARGIN * 2)
}
private var animationProgress = if (isOn) 1.0f else 0.0f
private var animationTimer: Timer? = null
private var sparkles = mutableListOf<Sparkle>()
private var sparkleTimer: Timer? = null
private data class Sparkle(
var x: Float,
var y: Float,
var size: Float,
var opacity: Float,
var life: Float,
val maxLife: Float,
val velocityX: Float,
val velocityY: Float,
val rotation: Float,
val rotationSpeed: Float
)
init {
preferredSize = Dimension(COMPONENT_WIDTH, COMPONENT_HEIGHT)
cursor = Cursor.getPredefinedCursor(Cursor.HAND_CURSOR)
addMouseListener(object : MouseAdapter() {
override fun mousePressed(e: MouseEvent?) {
e?.let { event ->
val toggleBounds = Rectangle(
SPARKLE_MARGIN,
SPARKLE_MARGIN,
TRACK_WIDTH,
TRACK_HEIGHT
)
if (toggleBounds.contains(event.point)) {
toggle()
}
}
}
})
}
fun setState(newState: Boolean, animate: Boolean = true) {
if (isOn != newState) {
isOn = newState
if (animate) {
animateToState()
} else {
animationProgress = if (isOn) 1.0f else 0.0f
repaint()
}
}
}
private fun toggle() {
isOn = !isOn
onToggle(isOn)
animateToState()
}
private fun animateToState() {
animationTimer?.stop()
val startProgress = animationProgress
val targetProgress = if (isOn) 1.0f else 0.0f
val startTime = System.currentTimeMillis()
animationTimer = Timer(TIMER_DELAY) { _ ->
val elapsed = System.currentTimeMillis() - startTime
val progress = (elapsed.toFloat() / ANIMATION_DURATION).coerceIn(0.0f, 1.0f)
animationProgress = startProgress + (targetProgress - startProgress) * progress
if (progress >= 1.0f) {
animationTimer?.stop()
animationProgress = targetProgress
if (isOn) {
triggerSparkles()
}
}
repaint()
}
animationTimer?.start()
}
private fun triggerSparkles() {
sparkles.clear()
sparkleTimer?.stop()
val trackX = SPARKLE_MARGIN.toFloat()
val trackY = SPARKLE_MARGIN.toFloat()
val thumbX = trackX + PADDING + 1.0f * (TRACK_WIDTH - THUMB_SIZE - 2 * PADDING)
val thumbY = trackY + PADDING
val thumbCenterX = thumbX + THUMB_SIZE / 2f
val thumbCenterY = thumbY + THUMB_SIZE / 2f
for (i in 0 until SPARKLE_COUNT) {
val angle = (i * 360f / SPARKLE_COUNT) * Math.PI / 180f
val maxSparkleSize = 5f
val safetyBuffer = 2f
val distanceToRightEdge = COMPONENT_WIDTH - thumbCenterX - maxSparkleSize - safetyBuffer
val distanceToBottomEdge = COMPONENT_HEIGHT - thumbCenterY - maxSparkleSize - safetyBuffer
val distanceToLeftEdge = thumbCenterX - maxSparkleSize - safetyBuffer
val distanceToTopEdge = thumbCenterY - maxSparkleSize - safetyBuffer
val maxSafeDistance =
minOf(distanceToRightEdge, distanceToBottomEdge, distanceToLeftEdge, distanceToTopEdge)
val constrainedMaxDistance = maxSafeDistance.coerceAtLeast(4f)
val distance = 4f + Math.random().toFloat() * (constrainedMaxDistance - 4f)
val sparkleX = thumbCenterX + cos(angle).toFloat() * distance
val sparkleY = thumbCenterY + sin(angle).toFloat() * distance
sparkles.add(
Sparkle(
x = sparkleX,
y = sparkleY,
size = 2f + Math.random().toFloat() * 3f,
opacity = 1f,
life = 0f,
maxLife = SPARKLE_DURATION.toFloat() + Math.random().toFloat() * 200f,
velocityX = (Math.random().toFloat() - 0.5f) * 0.5f,
velocityY = (Math.random().toFloat() - 0.5f) * 0.5f,
rotation = Math.random().toFloat() * 360f,
rotationSpeed = (Math.random().toFloat() - 0.5f) * 5f
)
)
}
startSparkleAnimation()
}
private fun startSparkleAnimation() {
sparkleTimer = Timer(TIMER_DELAY) { _ ->
var activeSparkles = false
for (sparkle in sparkles) {
sparkle.life += TIMER_DELAY
sparkle.x += sparkle.velocityX
sparkle.y += sparkle.velocityY
val lifeRatio = sparkle.life / sparkle.maxLife
sparkle.opacity = (1f - lifeRatio).coerceIn(0f, 1f)
sparkle.size = sparkle.size * 0.998f
if (sparkle.life < sparkle.maxLife) {
activeSparkles = true
}
}
if (!activeSparkles) {
sparkleTimer?.stop()
sparkles.clear()
}
repaint()
}
sparkleTimer?.start()
}
override fun paintComponent(g: Graphics) {
super.paintComponent(g)
val g2 = g.create() as Graphics2D
g2.setRenderingHint(RenderingHints.KEY_ANTIALIASING, RenderingHints.VALUE_ANTIALIAS_ON)
val trackX = SPARKLE_MARGIN.toFloat()
val trackY = SPARKLE_MARGIN.toFloat()
g2.color = if (isOn) Design.Colors.primary else Design.Colors.outline
g2.fill(createRoundRect(trackX, trackY, TRACK_WIDTH.toFloat(), TRACK_HEIGHT.toFloat(), TRACK_HEIGHT.toFloat()))
val thumbX = trackX + PADDING + animationProgress * (TRACK_WIDTH - THUMB_SIZE - 2 * PADDING)
val thumbY = trackY + PADDING
g2.color = Color(0, 0, 0, 20)
g2.fill(
createRoundRect(
thumbX + 1,
thumbY + 1,
THUMB_SIZE.toFloat(),
THUMB_SIZE.toFloat(),
THUMB_SIZE.toFloat()
)
)
g2.color = Color.WHITE
g2.fill(createRoundRect(thumbX, thumbY, THUMB_SIZE.toFloat(), THUMB_SIZE.toFloat(), THUMB_SIZE.toFloat()))
for (sparkle in sparkles) {
if (sparkle.opacity > 0) {
val alpha = (sparkle.opacity * 255).toInt().coerceIn(0, 255)
g2.composite = AlphaComposite.getInstance(AlphaComposite.SRC_OVER, sparkle.opacity)
val sparkleSize = sparkle.size
val cx = sparkle.x
val cy = sparkle.y
g2.color = Color(255, 215, 0, alpha)
g2.stroke = BasicStroke(1.5f, BasicStroke.CAP_ROUND, BasicStroke.JOIN_ROUND)
g2.drawLine(
(cx - sparkleSize).toInt(),
cy.toInt(),
(cx + sparkleSize).toInt(),
cy.toInt()
)
g2.drawLine(
cx.toInt(),
(cy - sparkleSize).toInt(),
cx.toInt(),
(cy + sparkleSize).toInt()
)
val diagSize = sparkleSize * 0.7f
g2.drawLine(
(cx - diagSize).toInt(),
(cy - diagSize).toInt(),
(cx + diagSize).toInt(),
(cy + diagSize).toInt()
)
g2.drawLine(
(cx - diagSize).toInt(),
(cy + diagSize).toInt(),
(cx + diagSize).toInt(),
(cy - diagSize).toInt()
)
}
}
g2.composite = AlphaComposite.getInstance(AlphaComposite.SRC_OVER, 1.0f)
g2.dispose()
}
override fun updateUI() {
super.updateUI()
repaint() // Repaint to use updated theme colors
}
private fun createRoundRect(
x: Float,
y: Float,
width: Float,
height: Float,
arcSize: Float
): RoundRectangle2D.Float {
return RoundRectangle2D.Float(x, y, width, height, arcSize, arcSize)
}
}
@@ -0,0 +1,479 @@
package net.portswigger.mcp.config
import burp.api.montoya.MontoyaApi
import burp.api.montoya.http.message.requests.HttpRequest
import burp.api.montoya.ui.editor.EditorOptions
import java.awt.*
import java.awt.event.ActionEvent
import java.awt.event.KeyEvent
import javax.swing.*
import javax.swing.border.EmptyBorder
object Dialogs {
private fun wrapText(text: String, maxWidth: Int = 50): String {
if (text.length <= maxWidth) return text
val words = text.split(" ")
val result = StringBuilder()
var currentLine = StringBuilder()
for (word in words) {
if (currentLine.length + word.length + 1 <= maxWidth) {
if (currentLine.isNotEmpty()) currentLine.append(" ")
currentLine.append(word)
} else {
if (result.isNotEmpty()) result.append("\n")
result.append(currentLine.toString())
currentLine = StringBuilder(word)
}
}
if (currentLine.isNotEmpty()) {
if (result.isNotEmpty()) result.append("\n")
result.append(currentLine.toString())
}
return result.toString()
}
private fun createDialog(parent: Component?): JDialog {
val parentWindow = SwingUtilities.getWindowAncestor(parent)
return JDialog(parentWindow, "", Dialog.ModalityType.APPLICATION_MODAL).apply {
background = Design.Colors.surface
defaultCloseOperation = JDialog.DISPOSE_ON_CLOSE
isResizable = true
val escapeAction = object : AbstractAction() {
override fun actionPerformed(e: ActionEvent?) {
dispose()
}
}
rootPane.getInputMap(JComponent.WHEN_IN_FOCUSED_WINDOW).put(
KeyStroke.getKeyStroke(KeyEvent.VK_ESCAPE, 0), "escape"
)
rootPane.actionMap.put("escape", escapeAction)
}
}
fun showMessageDialog(
parent: Component?, message: String, messageType: Int
) {
val dialog = createDialog(parent)
val iconLabel = when (messageType) {
JOptionPane.ERROR_MESSAGE -> JLabel("⚠").apply {
font = Design.Typography.headlineMedium
foreground = Design.Colors.error
horizontalAlignment = SwingConstants.CENTER
preferredSize = Dimension(40, 40)
}
JOptionPane.WARNING_MESSAGE -> JLabel("⚠").apply {
font = Design.Typography.headlineMedium
foreground = Design.Colors.warning
horizontalAlignment = SwingConstants.CENTER
preferredSize = Dimension(40, 40)
}
JOptionPane.INFORMATION_MESSAGE -> JLabel("ⓘ").apply {
font = Design.Typography.headlineMedium
foreground = Design.Colors.primary
horizontalAlignment = SwingConstants.CENTER
preferredSize = Dimension(40, 40)
}
else -> null
}
val messageLabel = JLabel(wrapText(message)).apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
horizontalAlignment = SwingConstants.CENTER
}
val contentPanel = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
background = Design.Colors.surface
border = EmptyBorder(Design.Spacing.XL, Design.Spacing.XL, Design.Spacing.LG, Design.Spacing.XL)
}
if (iconLabel != null) {
iconLabel.alignmentX = Component.CENTER_ALIGNMENT
contentPanel.add(iconLabel)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.MD))
}
messageLabel.alignmentX = Component.CENTER_ALIGNMENT
contentPanel.add(messageLabel)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.LG))
val okButton = Design.createFilledButton("OK").apply {
alignmentX = Component.CENTER_ALIGNMENT
addActionListener {
dialog.dispose()
}
}
contentPanel.add(okButton)
dialog.contentPane = contentPanel
dialog.pack()
dialog.setLocationRelativeTo(parent)
dialog.isVisible = true
}
fun showConfirmDialog(
parent: Component?, message: String, optionType: Int
): Int {
val dialog = createDialog(parent)
var result = JOptionPane.CANCEL_OPTION
val messageLabel = JLabel(message).apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
horizontalAlignment = SwingConstants.CENTER
}
val contentPanel = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
background = Design.Colors.surface
border = EmptyBorder(Design.Spacing.XL, Design.Spacing.XL, Design.Spacing.LG, Design.Spacing.XL)
}
messageLabel.alignmentX = Component.CENTER_ALIGNMENT
contentPanel.add(messageLabel)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.LG))
val buttonPanel = JPanel(FlowLayout(FlowLayout.CENTER, Design.Spacing.MD, 0)).apply {
background = Design.Colors.surface
alignmentX = Component.CENTER_ALIGNMENT
}
when (optionType) {
JOptionPane.YES_NO_OPTION -> {
val noButton = Design.createOutlinedButton("No").apply {
addActionListener {
result = JOptionPane.NO_OPTION
dialog.dispose()
}
}
val yesButton = Design.createFilledButton("Yes").apply {
addActionListener {
result = JOptionPane.YES_OPTION
dialog.dispose()
}
}
buttonPanel.add(noButton)
buttonPanel.add(yesButton)
}
JOptionPane.OK_CANCEL_OPTION -> {
val cancelButton = Design.createOutlinedButton("Cancel").apply {
addActionListener {
result = JOptionPane.CANCEL_OPTION
dialog.dispose()
}
}
val okButton = Design.createFilledButton("OK").apply {
addActionListener {
result = JOptionPane.OK_OPTION
dialog.dispose()
}
}
buttonPanel.add(cancelButton)
buttonPanel.add(okButton)
}
}
contentPanel.add(buttonPanel)
dialog.contentPane = contentPanel
dialog.pack()
dialog.setLocationRelativeTo(parent)
dialog.isVisible = true
return result
}
fun showInputDialog(
parent: Component?, message: String
): String? {
val dialog = createDialog(parent)
var result: String? = null
val messageLabel = JLabel(message).apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
}
val inputField = JTextField(20).apply {
font = Design.Typography.bodyLarge
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.outline, 1), BorderFactory.createEmptyBorder(
Design.Spacing.SM, Design.Spacing.MD, Design.Spacing.SM, Design.Spacing.MD
)
)
background = Design.Colors.listBackground
foreground = Design.Colors.onSurface
maximumSize = Dimension(Int.MAX_VALUE, preferredSize.height)
}
val contentPanel = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
background = Design.Colors.surface
border = EmptyBorder(Design.Spacing.XL, Design.Spacing.XL, Design.Spacing.LG, Design.Spacing.XL)
}
messageLabel.alignmentX = Component.LEFT_ALIGNMENT
contentPanel.add(messageLabel)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.MD))
inputField.alignmentX = Component.LEFT_ALIGNMENT
contentPanel.add(inputField)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.LG))
val buttonPanel = JPanel(FlowLayout(FlowLayout.RIGHT, Design.Spacing.MD, 0)).apply {
background = Design.Colors.surface
alignmentX = Component.LEFT_ALIGNMENT
}
val cancelButton = Design.createOutlinedButton("Cancel").apply {
addActionListener {
result = null
dialog.dispose()
}
}
val okButton = Design.createFilledButton("OK").apply {
addActionListener {
result = inputField.text?.takeIf { it.isNotBlank() }
dialog.dispose()
}
}
val enterAction = object : AbstractAction() {
override fun actionPerformed(e: ActionEvent?) {
result = inputField.text?.takeIf { it.isNotBlank() }
dialog.dispose()
}
}
inputField.getInputMap(JComponent.WHEN_FOCUSED).put(
KeyStroke.getKeyStroke(KeyEvent.VK_ENTER, 0), "enter"
)
inputField.actionMap.put("enter", enterAction)
buttonPanel.add(cancelButton)
buttonPanel.add(okButton)
contentPanel.add(buttonPanel)
dialog.contentPane = contentPanel
dialog.pack()
dialog.setLocationRelativeTo(parent)
SwingUtilities.invokeLater {
inputField.requestFocusInWindow()
}
dialog.isVisible = true
return result
}
fun showOptionDialog(
parent: Component?,
message: String,
options: Array<String>,
requestContent: String? = null,
api: MontoyaApi? = null
): Int {
val dialog = createDialog(parent)
var result = -1
val messageArea = JTextArea(message).apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
background = Design.Colors.surface
isEditable = false
isOpaque = false
lineWrap = true
wrapStyleWord = true
columns = 30
rows = 0
alignmentX = Component.CENTER_ALIGNMENT
}
val contentPanel = JPanel().apply {
background = Design.Colors.surface
}
if (!requestContent.isNullOrBlank()) {
contentPanel.layout = BorderLayout()
val leftPanel = JPanel().apply {
layout = BorderLayout()
background = Design.Colors.surface
minimumSize = Dimension(400, 300)
}
val requestComponent = if (api != null) {
try {
val httpRequestEditor = api.userInterface().createHttpRequestEditor(EditorOptions.READ_ONLY)
httpRequestEditor.request = HttpRequest.httpRequest(requestContent)
httpRequestEditor.uiComponent().apply {
minimumSize = Dimension(400, 200)
}
} catch (_: Exception) {
JTextArea(requestContent).apply {
font = Design.Typography.bodyMedium
foreground = Design.Colors.onSurface
background = Design.Colors.listBackground
isEditable = false
lineWrap = false
tabSize = 4
}.let { textArea ->
JScrollPane(textArea).apply {
verticalScrollBarPolicy = JScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JScrollPane.HORIZONTAL_SCROLLBAR_AS_NEEDED
border = BorderFactory.createLineBorder(Design.Colors.outline, 1)
minimumSize = Dimension(400, 200)
}
}
}
} else {
JTextArea(requestContent).apply {
font = Design.Typography.bodyMedium
foreground = Design.Colors.onSurface
background = Design.Colors.listBackground
isEditable = false
lineWrap = false
tabSize = 4
}.let { textArea ->
JScrollPane(textArea).apply {
verticalScrollBarPolicy = JScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JScrollPane.HORIZONTAL_SCROLLBAR_AS_NEEDED
border = BorderFactory.createLineBorder(Design.Colors.outline, 1)
minimumSize = Dimension(400, 200)
}
}
}
leftPanel.add(requestComponent, BorderLayout.CENTER)
val rightPanel = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
background = Design.Colors.surface
minimumSize = Dimension(400, 300)
maximumSize = Dimension(400, Int.MAX_VALUE)
preferredSize = Dimension(400, 400)
border = EmptyBorder(Design.Spacing.XL, Design.Spacing.LG, Design.Spacing.XL, Design.Spacing.XL)
}
messageArea.alignmentX = Component.CENTER_ALIGNMENT
rightPanel.add(messageArea)
rightPanel.add(Box.createVerticalStrut(Design.Spacing.LG))
rightPanel.add(Box.createVerticalGlue())
val buttonPanel = JPanel().apply {
layout = GridLayout(2, 2, Design.Spacing.SM, Design.Spacing.SM)
background = Design.Colors.surface
alignmentX = Component.CENTER_ALIGNMENT
preferredSize = Dimension(390, 80)
}
options.forEachIndexed { index, option ->
val button = when (index) {
0 -> Design.createFilledButton(option)
1, 2 -> Design.createOutlinedButton(option)
3 -> Design.createOutlinedButton(option).apply {
foreground = Design.Colors.error
border = BorderFactory.createLineBorder(Design.Colors.error, 1)
}
else -> Design.createOutlinedButton(option)
}.apply {
preferredSize = Dimension(190, 32)
font = font.deriveFont(10f)
addActionListener {
result = index
dialog.dispose()
}
}
buttonPanel.add(button)
}
rightPanel.add(buttonPanel)
contentPanel.add(leftPanel, BorderLayout.CENTER)
contentPanel.add(rightPanel, BorderLayout.EAST)
} else {
contentPanel.layout = BoxLayout(contentPanel, BoxLayout.Y_AXIS)
contentPanel.border =
EmptyBorder(Design.Spacing.XL, Design.Spacing.XL, Design.Spacing.XL, Design.Spacing.XL)
messageArea.alignmentX = Component.CENTER_ALIGNMENT
contentPanel.add(messageArea)
contentPanel.add(Box.createVerticalStrut(Design.Spacing.XL))
val buttonPanel = JPanel().apply {
layout = GridLayout(2, 2, Design.Spacing.SM, Design.Spacing.SM)
background = Design.Colors.surface
alignmentX = Component.CENTER_ALIGNMENT
preferredSize = Dimension(370, 80)
}
options.forEachIndexed { index, option ->
val button = when (index) {
0 -> Design.createFilledButton(option)
1, 2 -> Design.createOutlinedButton(option)
else -> Design.createTextButton(option)
}.apply {
preferredSize = Dimension(180, 32)
font = font.deriveFont(10f)
addActionListener {
result = index
dialog.dispose()
}
}
buttonPanel.add(button)
}
contentPanel.add(buttonPanel)
}
dialog.contentPane = contentPanel
if (requestContent.isNullOrBlank()) {
dialog.preferredSize = Dimension(420, 350)
} else {
dialog.preferredSize = Dimension(860, 400)
}
dialog.pack()
if (parent != null && parent.isDisplayable) {
dialog.setLocationRelativeTo(parent)
dialog.isAlwaysOnTop = true
dialog.toFront()
dialog.requestFocus()
} else {
val screenSize = Toolkit.getDefaultToolkit().screenSize
val dialogSize = dialog.size
dialog.setLocation(
(screenSize.width - dialogSize.width) / 2, (screenSize.height - dialogSize.height) / 2
)
}
dialog.isVisible = true
SwingUtilities.invokeLater {
dialog.isAlwaysOnTop = false
dialog.toFront()
}
return result
}
}
@@ -1,39 +1,164 @@
package net.portswigger.mcp.config
import burp.api.montoya.logging.Logging
import burp.api.montoya.persistence.PersistedObject
import java.lang.ref.WeakReference
import java.util.concurrent.CopyOnWriteArrayList
import kotlin.properties.ReadWriteProperty
import kotlin.reflect.KProperty
class McpConfig(storage: PersistedObject) {
class McpConfig(storage: PersistedObject, private val logging: Logging) {
var enabled by storage.boolean(true)
var configEditingTooling by storage.boolean(false)
var host by storage.string("127.0.0.1")
var port by storage.int(9876)
var requireHttpRequestApproval by storage.boolean(true)
var requireHistoryAccessApproval by storage.boolean(true)
private var _alwaysAllowHttpHistory by storage.boolean(false)
var alwaysAllowHttpHistory: Boolean
get() = _alwaysAllowHttpHistory
set(value) {
if (_alwaysAllowHttpHistory != value) {
_alwaysAllowHttpHistory = value
notifyHistoryAccessChanged()
}
}
private var _alwaysAllowWebSocketHistory by storage.boolean(false)
var alwaysAllowWebSocketHistory: Boolean
get() = _alwaysAllowWebSocketHistory
set(value) {
if (_alwaysAllowWebSocketHistory != value) {
_alwaysAllowWebSocketHistory = value
notifyHistoryAccessChanged()
}
}
private var _autoApproveTargets by storage.stringList("")
private val targetsChangeListeners = CopyOnWriteArrayList<ListenerRegistration>()
private val historyAccessChangeListeners = CopyOnWriteArrayList<ListenerRegistration>()
var autoApproveTargets: String
get() = _autoApproveTargets
set(value) {
if (_autoApproveTargets != value) {
_autoApproveTargets = value
notifyTargetsChanged()
}
}
fun addAutoApproveTarget(target: String): Boolean {
val currentTargets = getAutoApproveTargetsList()
if (target.trim().isNotEmpty() && !currentTargets.contains(target.trim())) {
val newTargets = currentTargets + target.trim()
autoApproveTargets = newTargets.joinToString(",")
return true
}
return false
}
fun removeAutoApproveTarget(target: String): Boolean {
val currentTargets = getAutoApproveTargetsList()
val newTargets = currentTargets.filter { it != target.trim() }
if (newTargets.size != currentTargets.size) {
autoApproveTargets = newTargets.joinToString(",")
return true
}
return false
}
fun getAutoApproveTargetsList(): List<String> {
return if (_autoApproveTargets.isBlank()) {
emptyList()
} else {
_autoApproveTargets.split(",").map { it.trim() }.filter { it.isNotEmpty() }
}
}
fun clearAutoApproveTargets() {
autoApproveTargets = ""
}
fun addTargetsChangeListener(listener: () -> Unit): ListenerHandle {
val registration = ListenerRegistration(listener)
targetsChangeListeners.add(registration)
return ListenerHandle { removeTargetsChangeListener(registration) }
}
private fun removeTargetsChangeListener(registration: ListenerRegistration) {
targetsChangeListeners.remove(registration)
}
private fun notifyTargetsChanged() {
cleanupStaleListeners(targetsChangeListeners)
val listeners = targetsChangeListeners.mapNotNull { it.listener.get() }
listeners.forEach { listener ->
try {
listener()
} catch (e: Exception) {
logging.logToError("Targets change listener failed: ${e.message}")
}
}
}
fun addHistoryAccessChangeListener(listener: () -> Unit): ListenerHandle {
val registration = ListenerRegistration(listener)
historyAccessChangeListeners.add(registration)
return ListenerHandle { removeHistoryAccessChangeListener(registration) }
}
private fun removeHistoryAccessChangeListener(registration: ListenerRegistration) {
historyAccessChangeListeners.remove(registration)
}
private fun notifyHistoryAccessChanged() {
cleanupStaleListeners(historyAccessChangeListeners)
val listeners = historyAccessChangeListeners.mapNotNull { it.listener.get() }
listeners.forEach { listener ->
try {
listener()
} catch (e: Exception) {
logging.logToError("History access change listener failed: ${e.message}")
}
}
}
private fun cleanupStaleListeners(listenerList: CopyOnWriteArrayList<ListenerRegistration>) {
val staleListeners = listenerList.filter { it.listener.get() == null }
listenerList.removeAll(staleListeners)
}
fun cleanup() {
targetsChangeListeners.clear()
historyAccessChangeListeners.clear()
}
}
fun PersistedObject.boolean(default: Boolean = false) =
PersistedDelegate(
getter = { key -> getBoolean(key) ?: default },
setter = { key, value -> setBoolean(key, value) }
)
PersistedDelegate(getter = { key -> getBoolean(key) ?: default }, setter = { key, value -> setBoolean(key, value) })
fun PersistedObject.string(default: String) =
PersistedDelegate(
getter = { key -> getString(key) ?: default },
setter = { key, value -> setString(key, value) }
)
PersistedDelegate(getter = { key -> getString(key) ?: default }, setter = { key, value -> setString(key, value) })
fun PersistedObject.int(default: Int) =
PersistedDelegate(
getter = { key -> getInteger(key) ?: default },
setter = { key, value -> setInteger(key, value) }
)
PersistedDelegate(getter = { key -> getInteger(key) ?: default }, setter = { key, value -> setInteger(key, value) })
fun PersistedObject.stringList(default: String) =
PersistedDelegate(getter = { key -> getString(key) ?: default }, setter = { key, value -> setString(key, value) })
class PersistedDelegate<T>(
private val getter: (name: String) -> T,
private val setter: (name: String, value: T) -> Unit
private val getter: (name: String) -> T, private val setter: (name: String, value: T) -> Unit
) : ReadWriteProperty<Any, T> {
override fun getValue(thisRef: Any, property: KProperty<*>) = getter(property.name)
override fun setValue(thisRef: Any, property: KProperty<*>, value: T) = setter(property.name, value)
}
class ListenerRegistration(listener: () -> Unit) {
val listener: WeakReference<() -> Unit> = WeakReference(listener)
}
fun interface ListenerHandle {
fun remove()
}
@@ -0,0 +1,41 @@
package net.portswigger.mcp.config
private const val MAX_TARGET_LENGTH = 255
object TargetValidation {
/**
* Validates whether a target string is in a valid format for auto-approve lists.
*
* Valid formats include:
* - Hostnames: example.com, localhost
* - IP addresses: 127.0.0.1, ::1
* - Hostnames with ports: example.com:8080, localhost:3000
* - IPv6 with ports: [::1]:8080
* - Wildcards: *.example.com, *.api.com
*
* @param target The target string to validate
* @return true if the target is valid, false otherwise
*/
fun isValidTarget(target: String): Boolean {
if (target.isBlank() || target.length > MAX_TARGET_LENGTH) return false
if (target.contains("\t") || target.contains("\n") || target.contains("\r")) return false
if (target.startsWith("[") && target.contains("]:")) {
val portPart = target.substringAfterLast(":")
val port = portPart.toIntOrNull()
return !(port == null || port < 1 || port > 65535)
}
val parts = target.split(":")
if (parts.size == 2) {
val port = parts[1].toIntOrNull()
if (port == null || port < 1 || port > 65535) return false
} else if (parts.size > 2) {
return true
}
return true
}
}
@@ -0,0 +1,111 @@
package net.portswigger.mcp.config.components
import net.portswigger.mcp.config.Design
import java.awt.Dimension
import java.awt.GridBagConstraints
import java.awt.GridBagLayout
import java.awt.Insets
import javax.swing.*
import javax.swing.Box.createVerticalStrut
import javax.swing.event.DocumentEvent
import javax.swing.event.DocumentListener
class AdvancedOptionsPanel(
private val hostField: JTextField,
private val portField: JTextField,
private val reinstallNotice: WarningLabel
) : JPanel() {
init {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
updateColors()
alignmentX = LEFT_ALIGNMENT
buildPanel()
setupFieldTracking()
}
override fun updateUI() {
super.updateUI()
updateColors()
}
private fun updateColors() {
background = Design.Colors.surface
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.outlineVariant, 1),
BorderFactory.createEmptyBorder(Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD)
)
}
private fun buildPanel() {
add(Design.createSectionLabel("Advanced Options"))
add(createVerticalStrut(Design.Spacing.MD))
val formPanel = createFormPanel(
"Server host:" to hostField, "Server port:" to portField
)
add(formPanel)
}
private fun setupFieldTracking() {
trackChanges(hostField)
trackChanges(portField)
}
private fun trackChanges(field: JTextField) {
field.document.addDocumentListener(object : DocumentListener {
override fun insertUpdate(e: DocumentEvent?) = handle()
override fun removeUpdate(e: DocumentEvent?) = handle()
override fun changedUpdate(e: DocumentEvent?) = handle()
fun handle() {
reinstallNotice.isVisible = true
}
})
}
private fun createFormPanel(vararg fields: Pair<String, JComponent>): JPanel {
val formPanel = JPanel(GridBagLayout()).apply {
isOpaque = false
alignmentX = LEFT_ALIGNMENT
}
val gbc = GridBagConstraints().apply {
insets = Insets(Design.Spacing.SM, 0, Design.Spacing.SM, Design.Spacing.MD)
anchor = GridBagConstraints.WEST
}
fields.forEachIndexed { index, (labelText, field) ->
gbc.gridx = 0
gbc.gridy = index
gbc.fill = GridBagConstraints.NONE
gbc.weightx = 0.0
formPanel.add(JLabel(labelText).apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
}, gbc)
gbc.gridx = 1
gbc.fill = GridBagConstraints.HORIZONTAL
gbc.weightx = 1.0
gbc.insets = Insets(Design.Spacing.SM, 0, Design.Spacing.SM, 0)
if (field is JTextField) {
field.preferredSize = Dimension(200, 32)
field.font = Design.Typography.bodyLarge
}
formPanel.add(field, gbc)
gbc.insets = Insets(Design.Spacing.SM, 0, Design.Spacing.SM, Design.Spacing.MD)
}
return formPanel
}
fun setFieldsEnabled(enabled: Boolean) {
hostField.isEnabled = enabled
portField.isEnabled = enabled
}
}
@@ -0,0 +1,299 @@
package net.portswigger.mcp.config.components
import net.portswigger.mcp.config.*
import net.portswigger.mcp.security.findBurpFrame
import java.awt.Component
import java.awt.Cursor
import java.awt.Dimension
import java.awt.FlowLayout
import java.awt.event.*
import javax.swing.*
import javax.swing.JOptionPane.*
class AutoApproveTargetsPanel(private val config: McpConfig) : JPanel() {
private var listenerHandle: ListenerHandle? = null
init {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
updateColors()
alignmentX = LEFT_ALIGNMENT
buildPanel()
}
override fun updateUI() {
super.updateUI()
updateColors()
}
private fun updateColors() {
background = Design.Colors.surface
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.outlineVariant, 1),
BorderFactory.createEmptyBorder(Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD)
)
}
private fun buildPanel() {
add(Design.createSectionLabel("Auto-Approved HTTP Targets"))
add(Box.createVerticalStrut(Design.Spacing.MD))
val descLabel = JLabel("Specify domains and hosts that can be accessed without approval.").apply {
alignmentX = LEFT_ALIGNMENT
font = Design.Typography.bodyMedium
foreground = Design.Colors.onSurfaceVariant
border = BorderFactory.createEmptyBorder(0, 0, Design.Spacing.SM, 0)
}
val examplesLabel = JLabel("Examples: example.com, localhost:8080, *.api.com").apply {
alignmentX = LEFT_ALIGNMENT
font = Design.Typography.labelMedium
foreground = Design.Colors.onSurfaceVariant
border = BorderFactory.createEmptyBorder(0, 0, Design.Spacing.MD, 0)
}
add(descLabel)
add(examplesLabel)
val listModel = DefaultListModel<String>()
val targetsList = createTargetsList(listModel)
updateTargetsList(listModel)
val refreshListener = {
SwingUtilities.invokeLater {
updateTargetsList(listModel)
}
}
listenerHandle = config.addTargetsChangeListener(refreshListener)
val scrollPane = createScrollPane(targetsList)
val tableContainer = createTableContainer(scrollPane)
add(tableContainer)
val buttonsPanel = createButtonsPanel(targetsList, listModel)
add(buttonsPanel)
}
private fun createTargetsList(listModel: DefaultListModel<String>): JList<String> {
return object : JList<String>(listModel) {
private var rolloverIndex = -1
init {
selectionMode = ListSelectionModel.SINGLE_SELECTION
visibleRowCount = 5
font = Design.Typography.bodyMedium
background = Design.Colors.listBackground
foreground = Design.Colors.onSurface
border = BorderFactory.createEmptyBorder(
Design.Spacing.SM, Design.Spacing.MD, Design.Spacing.SM, Design.Spacing.MD
)
cellRenderer = createCellRenderer()
addMouseMotionListener(createMouseMotionListener())
addMouseListener(createMouseListener())
addKeyListener(createKeyListener(listModel))
isFocusable = true
}
private fun createCellRenderer() = object : DefaultListCellRenderer() {
override fun getListCellRendererComponent(
list: JList<*>, value: Any?, index: Int, isSelected: Boolean, cellHasFocus: Boolean
): Component {
super.getListCellRendererComponent(list, value, index, isSelected, cellHasFocus)
border = BorderFactory.createEmptyBorder(
Design.Spacing.SM, Design.Spacing.MD, Design.Spacing.SM, Design.Spacing.MD
)
val isRollover = index == rolloverIndex && !isSelected
when {
isSelected -> {
background = Design.Colors.listSelectionBackground
foreground = Design.Colors.listSelectionForeground
}
isRollover -> {
background = Design.Colors.listHoverBackground
foreground = Design.Colors.onSurface
}
else -> {
background =
if (index % 2 == 0) Design.Colors.listBackground else Design.Colors.listAlternatingBackground
foreground = Design.Colors.onSurface
}
}
return this
}
}
private fun createMouseMotionListener() = object : MouseMotionAdapter() {
override fun mouseMoved(e: MouseEvent) {
try {
val index = locationToIndex(e.point)
val newRolloverIndex = if (index >= 0 && index < model.size && getCellBounds(
index, index
)?.contains(e.point) == true
) {
cursor = Cursor.getPredefinedCursor(Cursor.HAND_CURSOR)
index
} else {
cursor = Cursor.getDefaultCursor()
-1
}
if (rolloverIndex != newRolloverIndex) {
rolloverIndex = newRolloverIndex
repaint()
}
} catch (_: Exception) {
rolloverIndex = -1
cursor = Cursor.getDefaultCursor()
}
}
}
private fun createMouseListener() = object : MouseAdapter() {
override fun mouseExited(e: MouseEvent) {
if (rolloverIndex != -1) {
rolloverIndex = -1
cursor = Cursor.getDefaultCursor()
repaint()
}
}
}
private fun createKeyListener(listModel: DefaultListModel<String>) = object : KeyAdapter() {
override fun keyPressed(e: KeyEvent) {
when (e.keyCode) {
KeyEvent.VK_DELETE, KeyEvent.VK_BACK_SPACE -> {
if (selectedIndex >= 0 && selectedIndex < model.size) {
try {
removeTarget(selectedIndex, listModel)
e.consume()
} catch (ex: Exception) {
ex.printStackTrace()
}
}
}
}
}
}
}
}
private fun createScrollPane(targetsList: JList<String>): JScrollPane {
return JScrollPane(targetsList).apply {
val baseHeight = 220
val baseWidth = 400
val scaleFactor = Design.Spacing.MD / 16f
val responsiveHeight = (baseHeight * scaleFactor).toInt().coerceAtLeast(150)
val responsiveWidth = (baseWidth * scaleFactor).toInt().coerceAtLeast(250)
maximumSize = Dimension(Int.MAX_VALUE, responsiveHeight)
preferredSize = Dimension(responsiveWidth, responsiveHeight)
minimumSize = Dimension((responsiveWidth * 0.625f).toInt(), (responsiveHeight * 0.68f).toInt())
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.listBorder, 1), BorderFactory.createEmptyBorder(1, 1, 1, 1)
)
background = Design.Colors.listBackground
viewport.background = Design.Colors.listBackground
verticalScrollBarPolicy = JScrollPane.VERTICAL_SCROLLBAR_AS_NEEDED
horizontalScrollBarPolicy = JScrollPane.HORIZONTAL_SCROLLBAR_AS_NEEDED
}
}
private fun createTableContainer(scrollPane: JScrollPane): JPanel {
return JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
isOpaque = false
alignmentX = LEFT_ALIGNMENT
border = BorderFactory.createEmptyBorder(0, 0, Design.Spacing.MD, 0)
add(scrollPane)
}
}
private fun createButtonsPanel(targetsList: JList<String>, listModel: DefaultListModel<String>): JPanel {
val buttonsPanel = JPanel(FlowLayout(FlowLayout.LEFT, Design.Spacing.SM, Design.Spacing.SM)).apply {
isOpaque = false
alignmentX = LEFT_ALIGNMENT
border = BorderFactory.createEmptyBorder(Design.Spacing.SM, 0, 0, 0)
}
val addButton = Design.createFilledButton("Add").apply {
addActionListener {
val input = Dialogs.showInputDialog(
findBurpFrame(),
"Enter target (hostname or hostname:port):\nExamples: example.com, localhost:8080, *.api.com"
)
if (!input.isNullOrBlank()) {
val trimmed = input.trim()
if (TargetValidation.isValidTarget(trimmed)) {
addTarget(trimmed)
} else {
Dialogs.showMessageDialog(
findBurpFrame(),
"Invalid target format. Use hostname, IP address, hostname:port, or wildcard (*.domain)",
ERROR_MESSAGE
)
}
}
}
}
val removeButton = Design.createOutlinedButton("Remove").apply {
addActionListener {
val selectedIndex = targetsList.selectedIndex
if (selectedIndex >= 0) {
removeTarget(selectedIndex, listModel)
}
}
}
val clearButton = Design.createOutlinedButton("Clear All").apply {
addActionListener {
val result = Dialogs.showConfirmDialog(
findBurpFrame(), "Remove all auto-approved targets?", YES_NO_OPTION
)
if (result == YES_OPTION) {
clearAllTargets()
}
}
}
buttonsPanel.add(addButton)
buttonsPanel.add(removeButton)
buttonsPanel.add(clearButton)
return buttonsPanel
}
private fun updateTargetsList(listModel: DefaultListModel<String>) {
listModel.clear()
config.getAutoApproveTargetsList().forEach {
listModel.addElement(it)
}
}
private fun addTarget(target: String) {
config.addAutoApproveTarget(target)
}
private fun removeTarget(index: Int, listModel: DefaultListModel<String>) {
if (index >= 0 && index < listModel.size()) {
val target = listModel.getElementAt(index)
config.removeAutoApproveTarget(target)
}
}
private fun clearAllTargets() {
config.clearAutoApproveTargets()
}
fun cleanup() {
listenerHandle?.remove()
listenerHandle = null
}
}
@@ -0,0 +1,137 @@
package net.portswigger.mcp.config.components
import kotlinx.coroutines.CoroutineScope
import kotlinx.coroutines.Dispatchers
import kotlinx.coroutines.launch
import net.portswigger.mcp.Swing
import net.portswigger.mcp.config.Anchor
import net.portswigger.mcp.config.Design
import net.portswigger.mcp.config.Dialogs
import net.portswigger.mcp.config.McpConfig
import net.portswigger.mcp.providers.Provider
import java.awt.FlowLayout
import javax.swing.*
import javax.swing.Box.createVerticalStrut
import javax.swing.JOptionPane.*
import kotlin.concurrent.thread
class InstallationPanel(
private val config: McpConfig,
private val providers: List<Provider>,
private val reinstallNotice: WarningLabel,
private val parentComponent: JComponent
) : JPanel() {
init {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
updateColors()
alignmentX = LEFT_ALIGNMENT
buildPanel()
}
override fun updateUI() {
super.updateUI()
updateColors()
}
private fun updateColors() {
background = Design.Colors.surface
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.outlineVariant, 1),
BorderFactory.createEmptyBorder(Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD)
)
}
private fun buildPanel() {
add(Design.createSectionLabel("Installation"))
add(createVerticalStrut(Design.Spacing.SM))
val installOptions = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
alignmentX = LEFT_ALIGNMENT
isOpaque = false
}
val buttonRow = createButtonRow()
installOptions.add(buttonRow)
add(installOptions)
add(createVerticalStrut(Design.Spacing.SM))
val manualInstallPanel = createManualInstallPanel()
add(manualInstallPanel)
}
private fun createButtonRow(): JPanel {
val buttonRow = JPanel(FlowLayout(FlowLayout.LEFT, Design.Spacing.SM, Design.Spacing.SM)).apply {
alignmentX = LEFT_ALIGNMENT
isOpaque = false
}
providers.forEach { provider ->
val button = createProviderButton(provider)
buttonRow.add(button)
}
return buttonRow
}
private fun createProviderButton(provider: Provider): JButton {
return Design.createFilledButton(provider.installButtonText).apply {
addActionListener {
handleProviderInstall(provider)
}
}
}
private fun handleProviderInstall(provider: Provider) {
val confirmationText = provider.confirmationText
if (confirmationText != null) {
val result = Dialogs.showConfirmDialog(
parentComponent, confirmationText, YES_NO_OPTION
)
if (result != YES_OPTION) {
return
}
}
thread {
try {
val result = provider.install(config)
CoroutineScope(Dispatchers.Swing).launch {
reinstallNotice.isVisible = false
if (result != null) {
Dialogs.showMessageDialog(
parentComponent, result, INFORMATION_MESSAGE
)
}
}
} catch (e: Exception) {
CoroutineScope(Dispatchers.Swing).launch {
Dialogs.showMessageDialog(
parentComponent,
"Failed to install for ${provider.name}: ${e.message ?: e.javaClass.simpleName}",
ERROR_MESSAGE
)
}
}
}
}
private fun createManualInstallPanel(): JPanel {
return JPanel(FlowLayout(FlowLayout.LEFT, 0, 0)).apply {
alignmentX = LEFT_ALIGNMENT
isOpaque = false
add(
Anchor(
text = "Manual install steps",
url = "https://github.com/PortSwigger/mcp-server?tab=readme-ov-file#manual-installations"
)
)
}
}
}
@@ -0,0 +1,105 @@
package net.portswigger.mcp.config.components
import net.portswigger.mcp.config.Design
import java.awt.BorderLayout
import java.awt.GridBagConstraints
import java.awt.GridBagLayout
import javax.swing.BorderFactory
import javax.swing.BoxLayout
import javax.swing.JPanel
import javax.swing.JScrollPane
class ResponsiveColumnsPanel(private val leftPanel: JPanel, private val rightPanel: JScrollPane) : JPanel() {
private val minWidthForTwoColumns = 900
private val minWidthForLargePadding = 700
private var lastLayout = Layout.SINGLE_COLUMN
private var lastPaddingSize = PaddingSize.SMALL
private var isInitialized = false
enum class Layout { SINGLE_COLUMN, TWO_COLUMNS }
enum class PaddingSize { SMALL, LARGE }
init {
isInitialized = true
updateLayout()
}
override fun updateUI() {
super.updateUI()
if (isInitialized) {
updateLayout() // Reapply layout with updated theme colors
}
}
override fun doLayout() {
super.doLayout()
val currentLayout = if (width >= minWidthForTwoColumns) Layout.TWO_COLUMNS else Layout.SINGLE_COLUMN
val currentPaddingSize = if (width >= minWidthForLargePadding) PaddingSize.LARGE else PaddingSize.SMALL
if (currentLayout != lastLayout || currentPaddingSize != lastPaddingSize) {
lastLayout = currentLayout
lastPaddingSize = currentPaddingSize
updateLayout()
}
}
private fun updateLayout() {
removeAll()
val padding = when (lastPaddingSize) {
PaddingSize.LARGE -> Design.Spacing.LG
PaddingSize.SMALL -> Design.Spacing.SM
}
if (rightPanel.viewport.view is JPanel) {
val contentPanel = rightPanel.viewport.view as JPanel
contentPanel.border = BorderFactory.createEmptyBorder(padding, padding, padding, padding)
}
when (lastLayout) {
Layout.TWO_COLUMNS -> {
layout = GridBagLayout()
val c = GridBagConstraints().apply {
fill = GridBagConstraints.BOTH
weighty = 1.0
}
c.gridx = 0
c.gridy = 0
c.weightx = 0.35
add(leftPanel, c)
c.gridx = 1
c.weightx = 0.65
add(rightPanel, c)
}
Layout.SINGLE_COLUMN -> {
layout = BorderLayout()
val singleColumnPanel = JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
background = Design.Colors.surface
}
val headerWrapper = JPanel(BorderLayout()).apply {
isOpaque = false
border = BorderFactory.createEmptyBorder(padding, padding, Design.Spacing.MD, padding)
add(leftPanel, BorderLayout.CENTER)
}
singleColumnPanel.add(headerWrapper)
val scrollWrapper = JPanel(BorderLayout()).apply {
isOpaque = false
add(rightPanel, BorderLayout.CENTER)
}
singleColumnPanel.add(scrollWrapper)
add(singleColumnPanel, BorderLayout.CENTER)
}
}
revalidate()
repaint()
}
}
@@ -0,0 +1,185 @@
package net.portswigger.mcp.config.components
import net.portswigger.mcp.config.Design
import net.portswigger.mcp.config.McpConfig
import net.portswigger.mcp.config.ToggleSwitch
import java.awt.FlowLayout
import java.awt.event.ItemEvent
import javax.swing.*
import javax.swing.Box.createHorizontalStrut
import javax.swing.Box.createVerticalStrut
class ServerConfigurationPanel(
private val config: McpConfig,
private val enabledToggle: ToggleSwitch,
private val validationErrorLabel: WarningLabel
) : JPanel() {
private lateinit var alwaysAllowHttpHistoryCheckBox: JCheckBox
private lateinit var alwaysAllowWebSocketHistoryCheckBox: JCheckBox
init {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
updateColors()
alignmentX = LEFT_ALIGNMENT
buildPanel()
}
override fun updateUI() {
super.updateUI()
updateColors()
}
private fun updateColors() {
background = Design.Colors.surface
border = BorderFactory.createCompoundBorder(
BorderFactory.createLineBorder(Design.Colors.outlineVariant, 1),
BorderFactory.createEmptyBorder(Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD, Design.Spacing.MD)
)
}
private fun buildPanel() {
add(Design.createSectionLabel("Server Configuration"))
add(createVerticalStrut(Design.Spacing.MD))
val enabledPanel = createEnabledPanel()
add(enabledPanel)
add(createVerticalStrut(Design.Spacing.MD))
val configEditingToolingCheckBox = createCheckBoxWithSubtitle(
"Enable tools that can edit your config",
"WARNING: Can execute code",
config.configEditingTooling
) { config.configEditingTooling = it }
add(configEditingToolingCheckBox)
add(createVerticalStrut(Design.Spacing.MD))
val httpRequestApprovalCheckBox = createStandardCheckBox(
"Require approval for HTTP requests", config.requireHttpRequestApproval
) { config.requireHttpRequestApproval = it }
add(httpRequestApprovalCheckBox)
add(createVerticalStrut(Design.Spacing.MD))
val historyAccessApprovalCheckBox = createHistoryAccessApprovalCheckBox()
add(historyAccessApprovalCheckBox)
add(createVerticalStrut(Design.Spacing.SM))
alwaysAllowHttpHistoryCheckBox = createIndentedCheckBox(
"Always allow HTTP history access", config.alwaysAllowHttpHistory, config.requireHistoryAccessApproval
) { config.alwaysAllowHttpHistory = it }
add(alwaysAllowHttpHistoryCheckBox)
add(createVerticalStrut(Design.Spacing.SM))
alwaysAllowWebSocketHistoryCheckBox = createIndentedCheckBox(
"Always allow WebSocket history access",
config.alwaysAllowWebSocketHistory,
config.requireHistoryAccessApproval
) { config.alwaysAllowWebSocketHistory = it }
add(alwaysAllowWebSocketHistoryCheckBox)
add(validationErrorLabel)
}
private fun createEnabledPanel(): JPanel {
val enabledPanel = JPanel(FlowLayout(FlowLayout.LEFT, 0, 4)).apply {
isOpaque = false
alignmentX = LEFT_ALIGNMENT
}
enabledPanel.add(JLabel("Enabled").apply {
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
})
enabledPanel.add(createHorizontalStrut(Design.Spacing.MD))
enabledPanel.add(enabledToggle)
return enabledPanel
}
private fun createHistoryAccessApprovalCheckBox(): JCheckBox {
return createStandardCheckBox(
"Require approval for history access", config.requireHistoryAccessApproval
) { enabled ->
config.requireHistoryAccessApproval = enabled
if (!enabled) {
config.alwaysAllowHttpHistory = false
config.alwaysAllowWebSocketHistory = false
alwaysAllowHttpHistoryCheckBox.isSelected = false
alwaysAllowWebSocketHistoryCheckBox.isSelected = false
}
alwaysAllowHttpHistoryCheckBox.isEnabled = enabled
alwaysAllowWebSocketHistoryCheckBox.isEnabled = enabled
}
}
fun updateHistoryAccessCheckboxes() {
SwingUtilities.invokeLater {
alwaysAllowHttpHistoryCheckBox.isSelected = config.alwaysAllowHttpHistory
alwaysAllowWebSocketHistoryCheckBox.isSelected = config.alwaysAllowWebSocketHistory
}
}
private fun createStandardCheckBox(
text: String, initialValue: Boolean, onChange: (Boolean) -> Unit
): JCheckBox {
return JCheckBox(text).apply {
alignmentX = LEFT_ALIGNMENT
isSelected = initialValue
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
addItemListener { event ->
onChange(event.stateChange == ItemEvent.SELECTED)
}
}
}
private fun createIndentedCheckBox(
text: String, initialValue: Boolean, enabled: Boolean, onChange: (Boolean) -> Unit
): JCheckBox {
return JCheckBox(text).apply {
alignmentX = LEFT_ALIGNMENT
isSelected = initialValue
isEnabled = enabled
font = Design.Typography.bodyMedium
foreground = Design.Colors.onSurfaceVariant
border = BorderFactory.createEmptyBorder(0, Design.Spacing.LG, 0, 0)
addItemListener { event ->
onChange(event.stateChange == ItemEvent.SELECTED)
}
}
}
private fun createCheckBoxWithSubtitle(
mainText: String, subtitleText: String, initialValue: Boolean, onChange: (Boolean) -> Unit
): JPanel {
val checkBox = JCheckBox(mainText).apply {
alignmentX = LEFT_ALIGNMENT
isSelected = initialValue
font = Design.Typography.bodyLarge
foreground = Design.Colors.onSurface
addItemListener { event ->
onChange(event.stateChange == ItemEvent.SELECTED)
}
}
val subtitleLabel = JLabel(subtitleText).apply {
font = Design.Typography.labelMedium
foreground = Design.Colors.onSurfaceVariant
}
val subtitlePanel = JPanel(FlowLayout(FlowLayout.LEFT, 0, 0)).apply {
isOpaque = false
alignmentX = LEFT_ALIGNMENT
add(createHorizontalStrut(20))
add(subtitleLabel)
}
return JPanel().apply {
layout = BoxLayout(this, BoxLayout.Y_AXIS)
alignmentX = LEFT_ALIGNMENT
isOpaque = false
add(checkBox)
add(subtitlePanel)
}
}
}
@@ -0,0 +1,17 @@
package net.portswigger.mcp.config.components
import javax.swing.JLabel
import javax.swing.UIManager
class WarningLabel(content: String = "") : JLabel(content) {
init {
foreground = UIManager.getColor("Burp.warningBarBackground")
isVisible = false
alignmentX = LEFT_ALIGNMENT
}
override fun updateUI() {
super.updateUI()
foreground = UIManager.getColor("Burp.warningBarBackground")
}
}
@@ -8,7 +8,6 @@ import java.nio.file.Files
import java.nio.file.Path
import java.nio.file.StandardCopyOption
import javax.swing.JFileChooser
import javax.swing.JOptionPane
import kotlin.io.path.exists
import kotlin.io.path.readText
import kotlin.io.path.writeText
@@ -27,7 +26,8 @@ class ClaudeDesktopProvider(private val logging: Logging, private val proxyJarMa
override val name = "Claude Desktop"
override val installButtonText = "Install to $name"
override val confirmationText = "Install to $name?\nThis will create an entry within $name's MCP configuration file ($claudeConfigFileName)"
override val confirmationText =
"Install to $name?\nThis will create an entry within $name's MCP configuration file ($claudeConfigFileName)"
override fun install(config: McpConfig): String {
val proxyJarFile = proxyJarManager.getProxyJar()
@@ -117,7 +117,8 @@ class ClaudeDesktopProvider(private val logging: Logging, private val proxyJarMa
}
}
class ManualProxyInstallerProvider(private val logging: Logging, private val proxyJarManager: ProxyJarManager) : Provider {
class ManualProxyInstallerProvider(private val logging: Logging, private val proxyJarManager: ProxyJarManager) :
Provider {
override val name = "Proxy jar"
override val installButtonText = "Extract server proxy jar"
override val confirmationText = null
@@ -0,0 +1,89 @@
package net.portswigger.mcp.security
import net.portswigger.mcp.config.Dialogs
import net.portswigger.mcp.config.McpConfig
import javax.swing.SwingUtilities
import kotlin.coroutines.resume
import kotlin.coroutines.suspendCoroutine
enum class HistoryAccessType() {
HTTP_HISTORY(), WEBSOCKET_HISTORY();
}
interface HistoryAccessApprovalHandler {
suspend fun requestHistoryAccess(accessType: HistoryAccessType, config: McpConfig): Boolean
}
class SwingHistoryAccessApprovalHandler : HistoryAccessApprovalHandler {
override suspend fun requestHistoryAccess(
accessType: HistoryAccessType, config: McpConfig
): Boolean {
return suspendCoroutine { continuation ->
SwingUtilities.invokeLater {
val historyTypeName = when (accessType) {
HistoryAccessType.HTTP_HISTORY -> "HTTP history"
HistoryAccessType.WEBSOCKET_HISTORY -> "WebSocket history"
}
val message = buildString {
appendLine("An MCP client is requesting access to your Burp Suite $historyTypeName.")
appendLine()
appendLine("This may include sensitive data from previous web sessions.")
appendLine("Choose how you would like to respond:")
}
val options = arrayOf(
"Allow Once", "Always Allow $historyTypeName", "Deny"
)
val burpFrame = findBurpFrame()
val result = Dialogs.showOptionDialog(
burpFrame, message, options
)
when (result) {
0 -> {
continuation.resume(true)
}
1 -> {
when (accessType) {
HistoryAccessType.HTTP_HISTORY -> config.alwaysAllowHttpHistory = true
HistoryAccessType.WEBSOCKET_HISTORY -> config.alwaysAllowWebSocketHistory = true
}
continuation.resume(true)
}
else -> {
continuation.resume(false)
}
}
}
}
}
}
object HistoryAccessSecurity {
var approvalHandler: HistoryAccessApprovalHandler = SwingHistoryAccessApprovalHandler()
suspend fun checkHistoryAccessPermission(
accessType: HistoryAccessType, config: McpConfig
): Boolean {
if (!config.requireHistoryAccessApproval) {
return true
}
val isAlwaysAllowed = when (accessType) {
HistoryAccessType.HTTP_HISTORY -> config.alwaysAllowHttpHistory
HistoryAccessType.WEBSOCKET_HISTORY -> config.alwaysAllowWebSocketHistory
}
if (isAlwaysAllowed) {
return true
}
return approvalHandler.requestHistoryAccess(accessType, config)
}
}
@@ -0,0 +1,120 @@
package net.portswigger.mcp.security
import burp.api.montoya.MontoyaApi
import net.portswigger.mcp.config.Dialogs
import net.portswigger.mcp.config.McpConfig
import javax.swing.SwingUtilities
import kotlin.coroutines.resume
import kotlin.coroutines.suspendCoroutine
interface UserApprovalHandler {
suspend fun requestApproval(
hostname: String, port: Int, config: McpConfig, requestContent: String? = null, api: MontoyaApi? = null
): Boolean
}
class SwingUserApprovalHandler : UserApprovalHandler {
override suspend fun requestApproval(
hostname: String, port: Int, config: McpConfig, requestContent: String?, api: MontoyaApi?
): Boolean {
return suspendCoroutine { continuation ->
SwingUtilities.invokeLater {
val message = buildString {
appendLine("An MCP client is requesting to send an HTTP request to:")
appendLine()
appendLine("Target: $hostname:$port")
appendLine()
}
val options = arrayOf(
"Allow Once", "Always Allow Host", "Always Allow Host:Port", "Deny"
)
val burpFrame = findBurpFrame()
val result = Dialogs.showOptionDialog(
burpFrame, message, options, requestContent, api
)
when (result) {
0 -> {
continuation.resume(true)
}
1 -> {
config.addAutoApproveTarget(hostname)
continuation.resume(true)
}
2 -> {
config.addAutoApproveTarget("$hostname:$port")
continuation.resume(true)
}
else -> {
continuation.resume(false)
}
}
}
}
}
}
object HttpRequestSecurity {
var approvalHandler: UserApprovalHandler = SwingUserApprovalHandler()
private fun isAutoApproved(hostname: String, port: Int, config: McpConfig): Boolean {
val target = "$hostname:$port"
val hostOnly = hostname
val targets = config.getAutoApproveTargetsList()
return targets.any { approved ->
when {
approved.equals(target, ignoreCase = true) -> true
approved.equals(hostOnly, ignoreCase = true) -> true
approved.startsWith("*.") -> {
val domain = approved.substring(2)
isValidWildcardMatch(hostname, domain)
}
else -> false
}
}
}
private fun isValidWildcardMatch(hostname: String, domain: String): Boolean {
if (domain.isEmpty() || domain.contains("*")) return false
if (hostname.length <= domain.length) return false
val expectedSuffix = ".$domain"
if (!hostname.endsWith(expectedSuffix, ignoreCase = true)) return false
val subdomain = hostname.substring(0, hostname.length - expectedSuffix.length)
if (subdomain.isEmpty()) return false
return subdomain.split(".").all { label ->
label.isNotEmpty() && label.length <= 63 && !label.startsWith("-") && !label.endsWith("-") && label.matches(
Regex("^[a-zA-Z0-9-]+$")
)
}
}
suspend fun checkHttpRequestPermission(
hostname: String, port: Int, config: McpConfig, requestContent: String? = null, api: MontoyaApi? = null
): Boolean {
if (!config.requireHttpRequestApproval) {
return true
}
if (isAutoApproved(hostname, port, config)) {
return true
}
return approvalHandler.requestApproval(hostname, port, config, requestContent, api)
}
}
@@ -0,0 +1,20 @@
package net.portswigger.mcp.security
import java.awt.Frame
/**
* Finds the Burp Suite main frame or the largest available frame as fallback
*/
fun findBurpFrame(): Frame? {
val burpIdentifiers = listOf("Burp Suite", "Professional", "Community", "burp")
return Frame.getFrames().find { frame ->
frame.isVisible && frame.isDisplayable && burpIdentifiers.any { identifier ->
frame.title.contains(identifier, ignoreCase = true) ||
frame.javaClass.name.contains(identifier, ignoreCase = true) ||
frame.javaClass.simpleName.contains(identifier, ignoreCase = true)
}
} ?: Frame.getFrames()
.filter { it.isVisible && it.isDisplayable }
.maxByOrNull { it.width * it.height }
}
@@ -9,17 +9,51 @@ import burp.api.montoya.http.HttpService
import burp.api.montoya.http.message.HttpHeader
import burp.api.montoya.http.message.requests.HttpRequest
import io.modelcontextprotocol.kotlin.sdk.server.Server
import kotlinx.coroutines.runBlocking
import kotlinx.serialization.Serializable
import kotlinx.serialization.json.Json
import net.portswigger.mcp.config.McpConfig
import net.portswigger.mcp.schema.toSerializableForm
import net.portswigger.mcp.security.HistoryAccessSecurity
import net.portswigger.mcp.security.HistoryAccessType
import net.portswigger.mcp.security.HttpRequestSecurity
import java.awt.KeyboardFocusManager
import java.util.regex.Pattern
import javax.swing.JTextArea
private suspend fun checkHistoryPermissionOrDeny(
accessType: HistoryAccessType, config: McpConfig, api: MontoyaApi, logMessage: String
): Boolean {
val allowed = HistoryAccessSecurity.checkHistoryAccessPermission(accessType, config)
if (!allowed) {
api.logging().logToOutput("MCP $logMessage access denied")
return false
}
api.logging().logToOutput("MCP $logMessage access granted")
return true
}
private fun truncateIfNeeded(serialized: String): String {
return if (serialized.length > 5000) {
serialized.substring(0, 5000) + "... (truncated)"
} else {
serialized
}
}
fun Server.registerTools(api: MontoyaApi, config: McpConfig) {
mcpTool<SendHttp1Request>("Issues an HTTP/1.1 request and returns the response.") {
val allowed = runBlocking {
HttpRequestSecurity.checkHttpRequestPermission(targetHostname, targetPort, config, content, api)
}
if (!allowed) {
api.logging().logToOutput("MCP HTTP request denied: $targetHostname:$targetPort")
return@mcpTool "Send HTTP request denied by Burp Suite"
}
api.logging().logToOutput("MCP HTTP/1.1 request: $targetHostname:$targetPort")
val fixedContent = content.replace("\r", "").replace("\n", "\r\n")
val request = HttpRequest.httpRequest(toMontoyaService(), fixedContent)
@@ -29,6 +63,30 @@ fun Server.registerTools(api: MontoyaApi, config: McpConfig) {
}
mcpTool<SendHttp2Request>("Issues an HTTP/2 request and returns the response. Do NOT pass headers to the body parameter.") {
val http2RequestDisplay = buildString {
pseudoHeaders.forEach { (key, value) ->
val headerName = if (key.startsWith(":")) key else ":$key"
appendLine("$headerName: $value")
}
headers.forEach { (key, value) ->
appendLine("$key: $value")
}
if (requestBody.isNotBlank()) {
appendLine()
append(requestBody)
}
}
val allowed = runBlocking {
HttpRequestSecurity.checkHttpRequestPermission(targetHostname, targetPort, config, http2RequestDisplay, api)
}
if (!allowed) {
api.logging().logToOutput("MCP HTTP request denied: $targetHostname:$targetPort")
return@mcpTool "Send HTTP request denied by Burp Suite"
}
api.logging().logToOutput("MCP HTTP/2 request: $targetHostname:$targetPort")
val orderedPseudoHeaderNames = listOf(":scheme", ":method", ":path", ":authority")
val fixedPseudoHeaders = LinkedHashMap<String, String>().apply {
@@ -132,67 +190,52 @@ fun Server.registerTools(api: MontoyaApi, config: McpConfig) {
}
mcpPaginatedTool<GetProxyHttpHistory>("Displays items within the proxy HTTP history") {
api.proxy().history().asSequence()
.map {
// Limit the size of serialized data to prevent overflow
val serialized = Json.encodeToString(it.toSerializableForm())
if (serialized.length > 5000) {
// Truncate long responses to prevent chat overflow
val truncated = serialized.substring(0, 5000) + "... (truncated)"
truncated
} else {
serialized
}
}
val allowed = runBlocking {
checkHistoryPermissionOrDeny(HistoryAccessType.HTTP_HISTORY, config, api, "HTTP history")
}
if (!allowed) {
return@mcpPaginatedTool sequenceOf("HTTP history access denied by Burp Suite")
}
api.proxy().history().asSequence().map { truncateIfNeeded(Json.encodeToString(it.toSerializableForm())) }
}
mcpPaginatedTool<GetProxyHttpHistoryRegex>("Displays items matching a specified regex within the proxy HTTP history") {
val compiledRegex = Pattern.compile(regex)
val allowed = runBlocking {
checkHistoryPermissionOrDeny(HistoryAccessType.HTTP_HISTORY, config, api, "HTTP history")
}
if (!allowed) {
return@mcpPaginatedTool sequenceOf("HTTP history access denied by Burp Suite")
}
val compiledRegex = Pattern.compile(regex)
api.proxy().history { it.contains(compiledRegex) }.asSequence()
.map {
// Limit the size of serialized data to prevent overflow
val serialized = Json.encodeToString(it.toSerializableForm())
if (serialized.length > 5000) {
// Truncate long responses to prevent chat overflow
val truncated = serialized.substring(0, 5000) + "... (truncated)"
truncated
} else {
serialized
}
}
.map { truncateIfNeeded(Json.encodeToString(it.toSerializableForm())) }
}
mcpPaginatedTool<GetProxyWebsocketHistory>("Displays items within the proxy WebSocket history") {
val allowed = runBlocking {
checkHistoryPermissionOrDeny(HistoryAccessType.WEBSOCKET_HISTORY, config, api, "WebSocket history")
}
if (!allowed) {
return@mcpPaginatedTool sequenceOf("WebSocket history access denied by Burp Suite")
}
api.proxy().webSocketHistory().asSequence()
.map {
// Limit the size of serialized data to prevent overflow
val serialized = Json.encodeToString(it.toSerializableForm())
if (serialized.length > 5000) {
// Truncate long responses to prevent chat overflow
val truncated = serialized.substring(0, 5000) + "... (truncated)"
truncated
} else {
serialized
}
}
.map { truncateIfNeeded(Json.encodeToString(it.toSerializableForm())) }
}
mcpPaginatedTool<GetProxyWebsocketHistoryRegex>("Displays items matching a specified regex within the proxy WebSocket history") {
val compiledRegex = Pattern.compile(regex)
val allowed = runBlocking {
checkHistoryPermissionOrDeny(HistoryAccessType.WEBSOCKET_HISTORY, config, api, "WebSocket history")
}
if (!allowed) {
return@mcpPaginatedTool sequenceOf("WebSocket history access denied by Burp Suite")
}
val compiledRegex = Pattern.compile(regex)
api.proxy().webSocketHistory { it.contains(compiledRegex) }.asSequence()
.map {
// Limit the size of serialized data to prevent overflow
val serialized = Json.encodeToString(it.toSerializableForm())
if (serialized.length > 5000) {
// Truncate long responses to prevent chat overflow
val truncated = serialized.substring(0, 5000) + "... (truncated)"
truncated
} else {
serialized
}
}
.map { truncateIfNeeded(Json.encodeToString(it.toSerializableForm())) }
}
mcpTool<SetTaskExecutionEngineState>("Sets the state of Burp's task execution engine (paused or unpaused)") {
@@ -234,8 +277,7 @@ fun getActiveEditor(api: MontoyaApi): JTextArea? {
val focusManager = KeyboardFocusManager.getCurrentKeyboardFocusManager()
val permanentFocusOwner = focusManager.permanentFocusOwner
val isInBurpWindow = generateSequence(permanentFocusOwner) { it.parent }
.any { it == frame }
val isInBurpWindow = generateSequence(permanentFocusOwner) { it.parent }.any { it == frame }
return if (isInBurpWindow && permanentFocusOwner is JTextArea) {
permanentFocusOwner
@@ -1,6 +1,7 @@
package net.portswigger.mcp
import burp.api.montoya.MontoyaApi
import burp.api.montoya.logging.Logging
import burp.api.montoya.persistence.PersistedObject
import io.mockk.every
import io.mockk.mockk
@@ -31,12 +32,16 @@ class McpServerIntegrationTest {
every { persistedObject.setInteger(any(), any()) } returns Unit
}
private val config = McpConfig(persistedObject)
private val mockLogging = mockk<Logging>().apply {
every { logToError(any<String>()) } returns Unit
every { logToOutput(any<String>()) } returns Unit
}
private val config = McpConfig(persistedObject, mockLogging)
@BeforeEach
fun setup() {
serverManager.start(config) { state ->
println("Server state changed: $state")
if (state is ServerState.Running) {
serverStarted = true
}
@@ -0,0 +1,252 @@
package net.portswigger.mcp.config
import burp.api.montoya.logging.Logging
import burp.api.montoya.persistence.PersistedObject
import io.mockk.every
import io.mockk.mockk
import io.mockk.verify
import org.junit.jupiter.api.Assertions.*
import org.junit.jupiter.api.BeforeEach
import org.junit.jupiter.api.Test
class McpConfigTest {
private lateinit var persistedObject: PersistedObject
private lateinit var config: McpConfig
private lateinit var mockLogging: Logging
@BeforeEach
fun setup() {
val storage = mutableMapOf<String, Any>()
persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } answers {
val key = firstArg<String>()
storage[key] as? Boolean ?: when (key) {
"enabled" -> true
"requireHttpRequestApproval" -> true
else -> false
}
}
every { getString(any()) } answers { storage[firstArg()] as? String ?: "" }
every { getInteger(any()) } answers { storage[firstArg()] as? Int ?: 0 }
every { setBoolean(any(), any()) } answers {
storage[firstArg()] = secondArg<Boolean>()
}
every { setString(any(), any()) } answers {
storage[firstArg()] = secondArg<String>()
}
every { setInteger(any(), any()) } answers {
storage[firstArg()] = secondArg<Int>()
}
}
mockLogging = mockk<Logging>().apply {
every { logToError(any<String>()) } returns Unit
}
config = McpConfig(persistedObject, mockLogging)
}
@Test
fun `addAutoApproveTarget should add new target`() {
val result = config.addAutoApproveTarget("example.com")
assertTrue(result)
assertEquals("example.com", config.autoApproveTargets)
verify { persistedObject.setString("_autoApproveTargets", "example.com") }
}
@Test
fun `addAutoApproveTarget should not add duplicate target`() {
config.addAutoApproveTarget("example.com")
val result = config.addAutoApproveTarget("example.com")
assertFalse(result)
assertEquals("example.com", config.autoApproveTargets)
}
@Test
fun `addAutoApproveTarget should trim whitespace`() {
val result = config.addAutoApproveTarget(" example.com ")
assertTrue(result)
assertEquals("example.com", config.autoApproveTargets)
}
@Test
fun `addAutoApproveTarget should not add empty target`() {
val result = config.addAutoApproveTarget(" ")
assertFalse(result)
assertEquals("", config.autoApproveTargets)
}
@Test
fun `addAutoApproveTarget should handle multiple targets`() {
config.addAutoApproveTarget("example.com")
config.addAutoApproveTarget("test.org")
assertEquals("example.com,test.org", config.autoApproveTargets)
assertEquals(listOf("example.com", "test.org"), config.getAutoApproveTargetsList())
}
@Test
fun `removeAutoApproveTarget should remove existing target`() {
config.addAutoApproveTarget("example.com")
config.addAutoApproveTarget("test.org")
val result = config.removeAutoApproveTarget("example.com")
assertTrue(result)
assertEquals("test.org", config.autoApproveTargets)
assertEquals(listOf("test.org"), config.getAutoApproveTargetsList())
}
@Test
fun `removeAutoApproveTarget should return false for non-existing target`() {
config.addAutoApproveTarget("example.com")
val result = config.removeAutoApproveTarget("notfound.com")
assertFalse(result)
assertEquals("example.com", config.autoApproveTargets)
}
@Test
fun `clearAutoApproveTargets should remove all targets`() {
config.addAutoApproveTarget("example.com")
config.addAutoApproveTarget("test.org")
config.clearAutoApproveTargets()
assertEquals("", config.autoApproveTargets)
assertEquals(emptyList<String>(), config.getAutoApproveTargetsList())
}
@Test
fun `getAutoApproveTargetsList should handle empty config`() {
assertEquals(emptyList<String>(), config.getAutoApproveTargetsList())
}
@Test
fun `getAutoApproveTargetsList should parse comma-separated values`() {
val storage = mutableMapOf<String, Any>("_autoApproveTargets" to "example.com,test.org,*.api.com")
persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } answers { storage[firstArg()] as? Boolean ?: false }
every { getString(any()) } answers { storage[firstArg()] as? String ?: "" }
every { getInteger(any()) } answers { storage[firstArg()] as? Int ?: 0 }
every { setBoolean(any(), any()) } answers {
storage[firstArg()] = secondArg<Boolean>()
}
every { setString(any(), any()) } answers {
storage[firstArg()] = secondArg<String>()
}
every { setInteger(any(), any()) } answers {
storage[firstArg()] = secondArg<Int>()
}
}
config = McpConfig(persistedObject, mockLogging)
assertEquals(
listOf("example.com", "test.org", "*.api.com"), config.getAutoApproveTargetsList()
)
}
@Test
fun `getAutoApproveTargetsList should handle malformed input`() {
val storage = mutableMapOf<String, Any>("_autoApproveTargets" to "example.com,, ,test.org")
persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } answers { storage[firstArg()] as? Boolean ?: false }
every { getString(any()) } answers { storage[firstArg()] as? String ?: "" }
every { getInteger(any()) } answers { storage[firstArg()] as? Int ?: 0 }
every { setBoolean(any(), any()) } answers {
storage[firstArg()] = secondArg<Boolean>()
}
every { setString(any(), any()) } answers {
storage[firstArg()] = secondArg<String>()
}
every { setInteger(any(), any()) } answers {
storage[firstArg()] = secondArg<Int>()
}
}
config = McpConfig(persistedObject, mockLogging)
assertEquals(
listOf("example.com", "test.org"), config.getAutoApproveTargetsList()
)
}
@Test
fun `targets change listener should be notified`() {
var notificationCount = 0
val listener = {
notificationCount++
Unit
}
config.addTargetsChangeListener(listener)
config.addAutoApproveTarget("example.com")
assertEquals(1, notificationCount)
}
@Test
fun `targets change listener should handle exceptions`() {
val badListener = { throw RuntimeException("Test exception") }
val goodListener = { /* do nothing */ }
config.addTargetsChangeListener(badListener)
config.addTargetsChangeListener(goodListener)
assertDoesNotThrow {
config.addAutoApproveTarget("example.com")
}
}
@Test
fun `autoApproveTargets setter should only notify on actual changes`() {
var notificationCount = 0
val listener = {
notificationCount++
Unit
}
config.addTargetsChangeListener(listener)
config.autoApproveTargets = "example.com"
assertEquals(1, notificationCount)
config.autoApproveTargets = "example.com"
assertEquals(1, notificationCount)
config.autoApproveTargets = "test.org"
assertEquals(2, notificationCount)
}
@Test
fun `configEditingTooling should persist correctly`() {
assertFalse(config.configEditingTooling)
config.configEditingTooling = true
assertTrue(config.configEditingTooling)
verify { persistedObject.setBoolean("configEditingTooling", true) }
config.configEditingTooling = false
assertFalse(config.configEditingTooling)
verify { persistedObject.setBoolean("configEditingTooling", false) }
}
@Test
fun `requireHttpRequestApproval should persist correctly`() {
assertTrue(config.requireHttpRequestApproval)
config.requireHttpRequestApproval = false
assertFalse(config.requireHttpRequestApproval)
verify { persistedObject.setBoolean("requireHttpRequestApproval", false) }
config.requireHttpRequestApproval = true
assertTrue(config.requireHttpRequestApproval)
verify { persistedObject.setBoolean("requireHttpRequestApproval", true) }
}
}
@@ -0,0 +1,65 @@
package net.portswigger.mcp.config
import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.Test
class TargetValidationTest {
private fun isValidTarget(target: String): Boolean {
return TargetValidation.isValidTarget(target)
}
@Test
fun `isValidTarget should accept valid formats`() {
// Basic hostnames
assertTrue(isValidTarget("example.com"))
assertTrue(isValidTarget("test.org"))
assertTrue(isValidTarget("sub.domain.co.uk"))
assertTrue(isValidTarget("localhost"))
assertTrue(isValidTarget("127.0.0.1"))
// With ports
assertTrue(isValidTarget("example.com:80"))
assertTrue(isValidTarget("example.com:8080"))
assertTrue(isValidTarget("localhost:3000"))
assertTrue(isValidTarget("127.0.0.1:9876"))
// Wildcards
assertTrue(isValidTarget("*.example.com"))
assertTrue(isValidTarget("*.api.test.org"))
assertTrue(isValidTarget("*.co.uk"))
// IPv6 formats
assertTrue(isValidTarget("::1"))
assertTrue(isValidTarget("[::1]:8080"))
// Edge cases with permissive validation
assertTrue(isValidTarget("256.0.0.1")) // Invalid IPv4 but allowed
assertTrue(isValidTarget("test@example.com")) // Special chars allowed
assertTrue(isValidTarget("*.*.com")) // Multiple wildcards allowed
}
@Test
fun `isValidTarget should reject invalid formats`() {
// Empty/blank input
assertFalse(isValidTarget(""))
assertFalse(isValidTarget(" "))
// Invalid ports
assertFalse(isValidTarget("example.com:"))
assertFalse(isValidTarget("example.com:abc"))
assertFalse(isValidTarget("example.com:0"))
assertFalse(isValidTarget("example.com:65536"))
// Control characters
assertFalse(isValidTarget("example\tcom"))
assertFalse(isValidTarget("example\ncom"))
assertFalse(isValidTarget("example\rcom"))
// Oversized input
assertFalse(isValidTarget("a".repeat(256)))
}
}
@@ -0,0 +1,263 @@
package net.portswigger.mcp.security
import burp.api.montoya.logging.Logging
import burp.api.montoya.persistence.PersistedObject
import io.mockk.coEvery
import io.mockk.every
import io.mockk.mockk
import kotlinx.coroutines.runBlocking
import net.portswigger.mcp.config.McpConfig
import org.junit.jupiter.api.AfterEach
import org.junit.jupiter.api.Assertions.assertFalse
import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.BeforeEach
import org.junit.jupiter.api.Test
class HttpRequestSecurityTest {
private lateinit var persistedObject: PersistedObject
private lateinit var config: McpConfig
private lateinit var mockApprovalHandler: UserApprovalHandler
private lateinit var originalApprovalHandler: UserApprovalHandler
private lateinit var mockLogging: Logging
@BeforeEach
fun setup() {
originalApprovalHandler = HttpRequestSecurity.approvalHandler
mockApprovalHandler = mockk<UserApprovalHandler>()
HttpRequestSecurity.approvalHandler = mockApprovalHandler
val storage = mutableMapOf<String, Any>(
"enabled" to true,
"configEditingTooling" to false,
"requireHttpRequestApproval" to true,
"host" to "127.0.0.1",
"_autoApproveTargets" to "",
"port" to 9876
)
persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } answers { storage[firstArg()] as? Boolean ?: false }
every { getString(any()) } answers { storage[firstArg()] as? String ?: "" }
every { getInteger(any()) } answers { storage[firstArg()] as? Int ?: 0 }
every { setBoolean(any(), any()) } answers {
storage[firstArg()] = secondArg<Boolean>()
}
every { setString(any(), any()) } answers {
storage[firstArg()] = secondArg<String>()
}
every { setInteger(any(), any()) } answers {
storage[firstArg()] = secondArg<Int>()
}
}
mockLogging = mockk<Logging>().apply {
every { logToError(any<String>()) } returns Unit
}
config = McpConfig(persistedObject, mockLogging)
}
@AfterEach
fun tearDown() {
HttpRequestSecurity.approvalHandler = originalApprovalHandler
}
@Test
fun `checkHttpRequestPermission should allow when approval disabled`() {
config.requireHttpRequestApproval = false
runBlocking {
val result = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertTrue(result)
}
}
@Test
fun `checkHttpRequestPermission should allow auto-approved hostname`() {
config.addAutoApproveTarget("example.com")
runBlocking {
val result = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertTrue(result)
}
}
@Test
fun `checkHttpRequestPermission should allow auto-approved hostname with port`() {
config.addAutoApproveTarget("example.com:8080")
coEvery { mockApprovalHandler.requestApproval("example.com", 80, config, any()) } returns false
runBlocking {
val result1 = HttpRequestSecurity.checkHttpRequestPermission("example.com", 8080, config)
assertTrue(result1)
val result2 = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertFalse(result2)
}
}
@Test
fun `checkHttpRequestPermission should allow wildcard domains`() {
config.addAutoApproveTarget("*.example.com")
coEvery { mockApprovalHandler.requestApproval("example.com", 80, config) } returns false
runBlocking {
val result1 = HttpRequestSecurity.checkHttpRequestPermission("api.example.com", 80, config)
assertTrue(result1)
val result2 = HttpRequestSecurity.checkHttpRequestPermission("test.example.com", 443, config)
assertTrue(result2)
val result3 = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertFalse(result3)
}
}
@Test
fun `checkHttpRequestPermission should be case insensitive`() {
config.addAutoApproveTarget("Example.COM")
runBlocking {
val result1 = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertTrue(result1)
val result2 = HttpRequestSecurity.checkHttpRequestPermission("EXAMPLE.COM", 80, config)
assertTrue(result2)
}
}
@Test
fun `checkHttpRequestPermission should handle multiple targets`() {
config.addAutoApproveTarget("example.com")
config.addAutoApproveTarget("test.org:8080")
config.addAutoApproveTarget("*.api.com")
coEvery { mockApprovalHandler.requestApproval("test.org", 80, config, any()) } returns false
coEvery { mockApprovalHandler.requestApproval("notfound.com", 80, config, any()) } returns false
runBlocking {
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("example.com", 443, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("test.org", 8080, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("v1.api.com", 443, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("test.org", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("notfound.com", 80, config))
}
}
@Test
fun `checkHttpRequestPermission should handle empty auto-approve list`() {
coEvery { mockApprovalHandler.requestApproval("example.com", 80, config) } returns false
runBlocking {
val result = HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config)
assertFalse(result)
}
}
@Test
fun `isAutoApproved should handle malformed targets gracefully`() {
val storage = mutableMapOf<String, Any>(
"enabled" to true,
"configEditingTooling" to false,
"requireHttpRequestApproval" to true,
"host" to "127.0.0.1",
"_autoApproveTargets" to "example.com,, ,test.org",
"port" to 9876
)
persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } answers { storage[firstArg()] as? Boolean ?: false }
every { getString(any()) } answers { storage[firstArg()] as? String ?: "" }
every { getInteger(any()) } answers { storage[firstArg()] as? Int ?: 0 }
every { setBoolean(any(), any()) } answers {
storage[firstArg()] = secondArg<Boolean>()
}
every { setString(any(), any()) } answers {
storage[firstArg()] = secondArg<String>()
}
every { setInteger(any(), any()) } answers {
storage[firstArg()] = secondArg<Int>()
}
}
config = McpConfig(persistedObject, mockLogging)
coEvery { mockApprovalHandler.requestApproval("empty.com", 80, config, any()) } returns false
runBlocking {
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("example.com", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("test.org", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("empty.com", 80, config))
}
}
@Test
fun `wildcard matching should be secure and comprehensive`() {
config.addAutoApproveTarget("*.example.com")
config.addAutoApproveTarget("*.example.org")
coEvery { mockApprovalHandler.requestApproval(any(), any(), config, any()) } returns false
runBlocking {
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("api.example.com", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("test.example.com", 443, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("a.example.com", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("maliciousexample.com", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("notexample.com", 80, config))
assertFalse(
HttpRequestSecurity.checkHttpRequestPermission(
"example.com", 80, config
)
)
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("api.example.org", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("v1.api.example.org", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("test.example.net", 80, config))
}
}
@Test
fun `broad wildcards should work correctly`() {
config.addAutoApproveTarget("*.com")
coEvery { mockApprovalHandler.requestApproval(any(), any(), config, any()) } returns false
runBlocking {
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("anything.com", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("test.com", 443, config))
assertTrue(
HttpRequestSecurity.checkHttpRequestPermission(
"maliciousexample.com", 80, config
)
)
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("com", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("test.org", 80, config))
}
}
@Test
fun `invalid targets in list should not break matching`() {
config.addAutoApproveTarget("256.0.0.1") // Invalid IPv4
config.addAutoApproveTarget("test@domain.com") // Special chars
config.addAutoApproveTarget("*.*.com") // Multiple wildcards
config.addAutoApproveTarget("example..com") // Double dots
config.addAutoApproveTarget("valid.com") // Valid target
coEvery { mockApprovalHandler.requestApproval(any(), any(), config, any()) } returns false
runBlocking {
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("valid.com", 80, config))
assertTrue(HttpRequestSecurity.checkHttpRequestPermission("256.0.0.1", 80, config)) // Exact match works
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("other.com", 80, config))
assertFalse(HttpRequestSecurity.checkHttpRequestPermission("test.example.com", 80, config))
}
}
}
@@ -8,6 +8,7 @@ import burp.api.montoya.http.Http
import burp.api.montoya.http.HttpMode
import burp.api.montoya.http.message.HttpHeader
import burp.api.montoya.http.message.requests.HttpRequest
import burp.api.montoya.logging.Logging
import burp.api.montoya.persistence.PersistedObject
import burp.api.montoya.proxy.Proxy
import burp.api.montoya.proxy.ProxyHttpRequestResponse
@@ -49,14 +50,25 @@ class ToolsKtTest {
init {
val persistedObject = mockk<PersistedObject>().apply {
every { getBoolean(any()) } returns true
every { getString(any()) } returns "127.0.0.1"
every { getBoolean("enabled") } returns true
every { getBoolean("configEditingTooling") } returns true
every { getBoolean("requireHttpRequestApproval") } returns false
every { getBoolean("requireHistoryAccessApproval") } returns false
every { getBoolean("_alwaysAllowHttpHistory") } returns false
every { getBoolean("_alwaysAllowWebSocketHistory") } returns false
every { getString("host") } returns "127.0.0.1"
every { getString("autoApproveTargets") } returns ""
every { getInteger("port") } returns testPort
every { setBoolean(any(), any()) } returns Unit
every { setString(any(), any()) } returns Unit
every { setInteger(any(), any()) } returns Unit
}
config = McpConfig(persistedObject)
val mockLogging = mockk<Logging>().apply {
every { logToError(any<String>()) } returns Unit
every { logToOutput(any<String>()) } returns Unit
}
config = McpConfig(persistedObject, mockLogging)
mockkStatic(HttpHeader::class)
mockkStatic(burp.api.montoya.http.HttpService::class)