diff --git a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevServerHelper.kt b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevServerHelper.kt index b9eb8ad948b3..8533b6035573 100644 --- a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevServerHelper.kt +++ b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevServerHelper.kt @@ -16,6 +16,7 @@ import android.content.Context import android.net.Uri import android.os.AsyncTask import android.provider.Settings.Secure +import androidx.annotation.VisibleForTesting import com.facebook.common.logging.FLog import com.facebook.react.bridge.ReactContext import com.facebook.react.common.ReactConstants @@ -39,6 +40,10 @@ import java.io.UnsupportedEncodingException import java.security.MessageDigest import java.security.NoSuchAlgorithmException import java.util.Locale +import java.util.concurrent.Executor +import java.util.concurrent.LinkedBlockingQueue +import java.util.concurrent.ThreadPoolExecutor +import java.util.concurrent.TimeUnit import okhttp3.Call import okhttp3.Callback import okhttp3.OkHttpClient @@ -73,7 +78,8 @@ public open class DevServerHelper( public fun onPackagerDevMenuCommand() - // Allow apps to provide listeners for custom packager commands. + // Allow apps to provide listeners for custom packager commands. Handlers must not block, + // because closing the connection waits for a running handler. public fun customCommandHandlers(): Map? } @@ -87,7 +93,14 @@ public open class DevServerHelper( private val packagerStatusCheck: PackagerStatusCheck = PackagerStatusCheck(client) private val packageName: String = applicationContext.packageName + // Must stay single-threaded: open and close have to run in call order. + @VisibleForTesting + internal var packagerConnectionExecutor: Executor = + ThreadPoolExecutor(0, 1, 30, TimeUnit.SECONDS, LinkedBlockingQueue()) { runnable -> + Thread(runnable, PACKAGER_CONNECTION_THREAD_NAME) + } private var packagerClient: JSPackagerClient? = null + private var packagerClientHost: String? = null private var inspectorPackagerConnection: IInspectorPackagerConnection? = null /** Returns an opaque ID which is stable for the current combination of device and app, stable */ @@ -138,65 +151,62 @@ public open class DevServerHelper( get() = settings.isJSMinifyEnabled public fun openPackagerConnection(clientId: String?, commandListener: PackagerCommandListener) { - if (packagerClient != null) { - FLog.w(ReactConstants.TAG, "Packager connection already open, nooping.") - return - } - object : AsyncTask() { - @Deprecated("This needs to be rewritten to not use AsyncTasks") - override fun doInBackground(vararg backgroundParams: Void): Void? { - val handlers: MutableMap = mutableMapOf() - handlers["reload"] = - object : NotificationOnlyHandler() { - override fun onNotification(params: Any?) { - commandListener.onPackagerReloadCommand() - } - } - handlers["devMenu"] = - object : NotificationOnlyHandler() { - override fun onNotification(params: Any?) { - commandListener.onPackagerDevMenuCommand() - } - } - commandListener.customCommandHandlers()?.let { handlers.putAll(it) } - - val onPackagerConnectedCallback: ReconnectingWebSocket.ConnectionCallback = - object : ReconnectingWebSocket.ConnectionCallback { - override fun onConnected() { - commandListener.onPackagerConnected() - } - - override fun onDisconnected() { - commandListener.onPackagerDisconnected() - } - } - - checkNotNull(clientId) - packagerClient = - JSPackagerClient( - clientId, - packagerConnectionSettings, - handlers, - onPackagerConnectedCallback, - ) - .apply { init() } + val id = checkNotNull(clientId) + packagerConnectionExecutor.execute { + val host = packagerConnectionSettings.debugServerHost + packagerClient?.let { client -> + if (host == packagerClientHost) { + FLog.w(ReactConstants.TAG, "Packager connection already open, nooping.") + return@execute + } + // The dev server host changed without a close. + client.close() + packagerClient = null + } + val handlers: MutableMap = mutableMapOf() + handlers["reload"] = + object : NotificationOnlyHandler() { + override fun onNotification(params: Any?) { + commandListener.onPackagerReloadCommand() + } + } + handlers["devMenu"] = + object : NotificationOnlyHandler() { + override fun onNotification(params: Any?) { + commandListener.onPackagerDevMenuCommand() + } + } + commandListener.customCommandHandlers()?.let { handlers.putAll(it) } - return null + val onPackagerConnectedCallback: ReconnectingWebSocket.ConnectionCallback = + object : ReconnectingWebSocket.ConnectionCallback { + override fun onConnected() { + commandListener.onPackagerConnected() + } + + override fun onDisconnected() { + commandListener.onPackagerDisconnected() + } } - } - .executeOnExecutor(AsyncTask.THREAD_POOL_EXECUTOR) + + packagerClient = + JSPackagerClient( + id, + packagerConnectionSettings, + handlers, + onPackagerConnectedCallback, + ) + .apply { init() } + packagerClientHost = host + } } public fun closePackagerConnection() { - object : AsyncTask() { - @Deprecated("This class needs to be rewritten to don't use AsyncTasks") - override fun doInBackground(vararg params: Void): Void? { - packagerClient?.close() - packagerClient = null - return null - } - } - .executeOnExecutor(AsyncTask.THREAD_POOL_EXECUTOR) + packagerConnectionExecutor.execute { + packagerClient?.close() + packagerClient = null + packagerClientHost = null + } } public fun openInspectorConnection() { @@ -385,6 +395,7 @@ public open class DevServerHelper( private companion object { private const val DEBUGGER_MSG_DISABLE = "{ \"id\":1,\"method\":\"Debugger.disable\" }" + private const val PACKAGER_CONNECTION_THREAD_NAME = "ReactPackagerConnection" private fun getSHA256(string: String): String { val digest = diff --git a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevSupportManagerBase.kt b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevSupportManagerBase.kt index 8abdd78fad99..1c24fafeb065 100644 --- a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevSupportManagerBase.kt +++ b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/DevSupportManagerBase.kt @@ -52,6 +52,7 @@ import com.facebook.react.devsupport.DevServerHelper.PackagerCommandListener import com.facebook.react.devsupport.InspectorFlags.getFuseboxEnabled import com.facebook.react.devsupport.StackTraceHelper.convertJavaStackTrace import com.facebook.react.devsupport.StackTraceHelper.convertJsStackTrace +import com.facebook.react.devsupport.inspector.DevSupportHttpClient import com.facebook.react.devsupport.inspector.TracingState import com.facebook.react.devsupport.inspector.TracingStateProvider import com.facebook.react.devsupport.interfaces.BundleLoadCallback @@ -450,9 +451,24 @@ public abstract class DevSupportManagerBase( return@DevOptionHandler } - ChangeBundleLocationDialog.show(context, devSettings) { host: String -> - devSettings.packagerConnectionSettings.debugServerHost = host - handleReloadJS() + ChangeBundleLocationDialog.show(context, devSettings) { input: String -> + val host = input.trim() + // An empty host resets to the default. An invalid host would crash the connection. + if (host.isEmpty() || DevSupportHttpClient.isValidHost(host)) { + devSettings.packagerConnectionSettings.debugServerHost = host + devServerHelper.closePackagerConnection() + handleReloadJS() + } else { + Toast.makeText( + applicationContext, + applicationContext.getString( + R.string.catalyst_change_bundle_location_invalid, + host, + ), + Toast.LENGTH_LONG, + ) + .show() + } } } diff --git a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/inspector/DevSupportHttpClient.kt b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/inspector/DevSupportHttpClient.kt index 547c8d803a0b..ebbf41a75f08 100644 --- a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/inspector/DevSupportHttpClient.kt +++ b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/devsupport/inspector/DevSupportHttpClient.kt @@ -13,6 +13,7 @@ import com.facebook.react.modules.network.OkHttpClientProvider import java.util.concurrent.TimeUnit import okhttp3.ConnectionPool import okhttp3.Dispatcher +import okhttp3.HttpUrl.Companion.toHttpUrlOrNull import okhttp3.OkHttpClient /** @@ -68,4 +69,10 @@ internal object DevSupportHttpClient { * the host specifies port 443 explicitly (e.g. "example.com:443"). */ internal fun wsScheme(host: String): String = if (host.endsWith(":443")) "wss" else "ws" + + /** + * Returns whether OkHttp can build a request for the given host, for example "localhost:8081". + */ + internal fun isValidHost(host: String): Boolean = + "${httpScheme(host)}://$host/".toHttpUrlOrNull() != null } diff --git a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/packagerconnection/ReconnectingWebSocket.kt b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/packagerconnection/ReconnectingWebSocket.kt index ada82e7a8d30..897106e9316b 100644 --- a/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/packagerconnection/ReconnectingWebSocket.kt +++ b/packages/react-native/ReactAndroid/src/main/java/com/facebook/react/packagerconnection/ReconnectingWebSocket.kt @@ -77,16 +77,18 @@ public class ReconnectingWebSocket( } public fun closeQuietly() { - closed = true - closeWebSocketQuietly() - messageCallback = null + synchronized(this) { + closed = true + closeWebSocketQuietly() + messageCallback = null + } connectionCallback?.onDisconnected() } private fun closeWebSocketQuietly() { try { - webSocket?.close(1_000, "End of session") + webSocket?.close(CLOSE_NORMAL, CLOSE_REASON) } catch (e: Exception) { // swallow, no need to handle it here } @@ -100,6 +102,10 @@ public class ReconnectingWebSocket( @Synchronized override fun onOpen(webSocket: WebSocket, response: Response) { + if (closed) { + webSocket.close(CLOSE_NORMAL, CLOSE_REASON) + return + } this.webSocket = webSocket suppressConnectionErrors = false @@ -155,5 +161,7 @@ public class ReconnectingWebSocket( private val TAG: String = ReconnectingWebSocket::class.java.simpleName private const val RECONNECT_DELAY_MS = 2_000L + private const val CLOSE_NORMAL = 1_000 + private const val CLOSE_REASON = "End of session" } } diff --git a/packages/react-native/ReactAndroid/src/main/res/devsupport/values/strings.xml b/packages/react-native/ReactAndroid/src/main/res/devsupport/values/strings.xml index 9621bacc928f..e86a812a3fd7 100644 --- a/packages/react-native/ReactAndroid/src/main/res/devsupport/values/strings.xml +++ b/packages/react-native/ReactAndroid/src/main/res/devsupport/values/strings.xml @@ -8,6 +8,7 @@ You can connect either via USB (localhost - default) or Wifi. If you connect via USB and running with a physical device, make sure you:\n 1. Connect your device via USB\n 2. Set the bundle location to `localhost:8081`\n 3. Run this command in your terminal:\n      `%1$s` Apply Changes Cancel + Invalid bundler address: %1$s Failed to open DevTools. Please check that the dev server is running and reload the app. Open DevTools Finish performance trace diff --git a/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/DevServerHelperTest.kt b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/DevServerHelperTest.kt new file mode 100644 index 000000000000..8a2fcb06d91e --- /dev/null +++ b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/DevServerHelperTest.kt @@ -0,0 +1,175 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * + * This source code is licensed under the MIT license found in the + * LICENSE file in the root directory of this source tree. + */ + +package com.facebook.react.devsupport + +import android.net.Uri +import com.facebook.react.devsupport.DevServerHelper.PackagerCommandListener +import com.facebook.react.packagerconnection.PackagerConnectionSettings +import com.facebook.react.packagerconnection.ReconnectingWebSocket +import java.util.ArrayDeque +import java.util.concurrent.Executor +import java.util.concurrent.FutureTask +import java.util.concurrent.ThreadPoolExecutor +import java.util.concurrent.TimeUnit +import org.assertj.core.api.Assertions.assertThat +import org.assertj.core.api.Assertions.assertThatThrownBy +import org.junit.After +import org.junit.Before +import org.junit.Test +import org.junit.runner.RunWith +import org.mockito.MockedConstruction +import org.mockito.Mockito.mockConstruction +import org.mockito.kotlin.doAnswer +import org.mockito.kotlin.mock +import org.mockito.kotlin.verify +import org.mockito.kotlin.whenever +import org.robolectric.RobolectricTestRunner +import org.robolectric.RuntimeEnvironment + +@RunWith(RobolectricTestRunner::class) +class DevServerHelperTest { + private lateinit var helper: DevServerHelper + private lateinit var sockets: MockedConstruction + private val settings: PackagerConnectionSettings = mock() + private val peers = mutableListOf() + private val listener: PackagerCommandListener = mock() + private val pending = ArrayDeque() + private var onConnect: () -> Unit = {} + + @Before + fun setUp() { + whenever(settings.debugServerHost).thenReturn("127.0.0.1:8083") + whenever(settings.packageName).thenReturn("com.example.test") + sockets = mockSockets() + helper = DevServerHelper(mock(), RuntimeEnvironment.getApplication(), settings) + helper.packagerConnectionExecutor = Executor { pending.add(it) } + } + + @After + fun tearDown() { + helper.closePackagerConnection() + runPending() + sockets.close() + } + + @Test + fun openAndCloseAreQueued() { + helper.openPackagerConnection("test", listener) + helper.closePackagerConnection() + + assertThat(peers).isEmpty() + assertThat(pending).hasSize(2) + } + + @Test + fun defaultExecutorRunsTasksOnOneNamedThread() { + val executor = + DevServerHelper(mock(), RuntimeEnvironment.getApplication(), settings) + .packagerConnectionExecutor as ThreadPoolExecutor + assertThat(executor.maximumPoolSize).isEqualTo(1) + + val threadName = FutureTask { Thread.currentThread().name } + executor.execute(threadName) + assertThat(threadName.get(10, TimeUnit.SECONDS)).isEqualTo("ReactPackagerConnection") + } + + @Test + fun repeatedOpenCreatesOnlyOneConnection() { + helper.openPackagerConnection("test", listener) + helper.openPackagerConnection("test", listener) + runPending() + + assertThat(peers).hasSize(1) + assertThat(connectedPeers()).hasSize(1) + } + + @Test + fun closeQueuedAfterOpenRetiresTheClient() { + helper.openPackagerConnection("test", listener) + helper.closePackagerConnection() + runPending() + + assertThat(peers).hasSize(1) + assertThat(connectedPeers()).isEmpty() + } + + @Test + fun failedStartupThrowsAndAllowsRetry() { + val failure = IllegalStateException("WebSocket startup failed") + onConnect = { throw failure } + helper.openPackagerConnection("test", listener) + assertThatThrownBy { runPending() }.isSameAs(failure) + assertThat(peers).hasSize(1) + assertThat(connectedPeers()).isEmpty() + + onConnect = {} + helper.openPackagerConnection("test", listener) + runPending() + assertThat(peers).hasSize(2) + assertThat(connectedPeers()).containsExactly(peers.last()) + } + + @Test + fun openAfterHostChangeWithoutCloseUsesOnlyTheNewServer() { + helper.openPackagerConnection("test", listener) + runPending() + whenever(settings.debugServerHost).thenReturn("127.0.0.1:8082") + helper.openPackagerConnection("test", listener) + runPending() + + assertThat(peers).hasSize(2) + assertThat(connectedPeers().map { it.port }).containsExactly(8082) + connectedPeers().single().receiveReload() + verify(listener).onPackagerReloadCommand() + } + + private fun runPending() { + while (true) { + val task = pending.poll() ?: return + task.run() + } + } + + private fun mockSockets(): MockedConstruction = + mockConstruction(ReconnectingWebSocket::class.java) { socket, construction -> + val peer = + Peer( + construction.arguments()[0] as String, + construction.arguments()[1] as ReconnectingWebSocket.MessageCallback, + ) + val connectionCallback = + construction.arguments()[2] as ReconnectingWebSocket.ConnectionCallback + peers.add(peer) + doAnswer { + onConnect() + peer.connected = true + } + .whenever(socket) + .connect() + doAnswer { + peer.connected = false + connectionCallback.onDisconnected() + } + .whenever(socket) + .closeQuietly() + } + + private fun connectedPeers(): List = peers.filter { it.connected } + + private class Peer(val url: String, val callback: ReconnectingWebSocket.MessageCallback) { + @Volatile var connected = false + val port: Int + get() = Uri.parse(url).port + + fun receiveReload() { + callback.onMessage("""{"version":2,"method":"reload"}""") + } + + override fun toString(): String = url + } +} diff --git a/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/inspector/DevSupportHttpClientTest.kt b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/inspector/DevSupportHttpClientTest.kt new file mode 100644 index 000000000000..d3445da24e29 --- /dev/null +++ b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/devsupport/inspector/DevSupportHttpClientTest.kt @@ -0,0 +1,28 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * + * This source code is licensed under the MIT license found in the + * LICENSE file in the root directory of this source tree. + */ + +package com.facebook.react.devsupport.inspector + +import org.assertj.core.api.Assertions.assertThat +import org.junit.Test + +class DevSupportHttpClientTest { + @Test + fun acceptsHostsThatOkHttpCanConnectTo() { + assertThat(DevSupportHttpClient.isValidHost("localhost:8081")).isTrue() + assertThat(DevSupportHttpClient.isValidHost("10.0.2.2:8081")).isTrue() + assertThat(DevSupportHttpClient.isValidHost("example.com:443")).isTrue() + } + + @Test + fun rejectsHostsThatOkHttpCannotParse() { + assertThat(DevSupportHttpClient.isValidHost("localhost:80a")).isFalse() + assertThat(DevSupportHttpClient.isValidHost("foo bar:8081")).isFalse() + assertThat(DevSupportHttpClient.isValidHost("localhost:8081 ")).isFalse() + assertThat(DevSupportHttpClient.isValidHost("localhost:8081\t")).isFalse() + } +} diff --git a/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/packagerconnection/ReconnectingWebSocketTest.kt b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/packagerconnection/ReconnectingWebSocketTest.kt new file mode 100644 index 000000000000..0b509b0b2439 --- /dev/null +++ b/packages/react-native/ReactAndroid/src/test/java/com/facebook/react/packagerconnection/ReconnectingWebSocketTest.kt @@ -0,0 +1,117 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * + * This source code is licensed under the MIT license found in the + * LICENSE file in the root directory of this source tree. + */ + +package com.facebook.react.packagerconnection + +import com.facebook.react.packagerconnection.ReconnectingWebSocket.ConnectionCallback +import com.facebook.react.packagerconnection.ReconnectingWebSocket.MessageCallback +import java.util.concurrent.CopyOnWriteArrayList +import java.util.concurrent.CountDownLatch +import java.util.concurrent.TimeUnit +import okhttp3.Protocol +import okhttp3.Request +import okhttp3.Response +import okhttp3.WebSocket +import org.assertj.core.api.Assertions.assertThat +import org.junit.Test +import org.junit.runner.RunWith +import org.mockito.kotlin.any +import org.mockito.kotlin.doAnswer +import org.mockito.kotlin.eq +import org.mockito.kotlin.mock +import org.mockito.kotlin.never +import org.mockito.kotlin.verify +import org.mockito.kotlin.whenever +import org.robolectric.RobolectricTestRunner + +@RunWith(RobolectricTestRunner::class) +class ReconnectingWebSocketTest { + private val messageCallback: MessageCallback = mock() + private val connectionCallback: ConnectionCallback = mock() + private val socket = ReconnectingWebSocket(URL, messageCallback, connectionCallback) + + @Test + fun closeClosesAnOpenSocket() { + val webSocket: WebSocket = mock() + socket.onOpen(webSocket, response()) + + socket.closeQuietly() + + verify(webSocket).close(eq(1_000), any()) + } + + @Test + fun socketThatOpensAfterCloseIsClosedAndNotReported() { + val webSocket: WebSocket = mock() + socket.closeQuietly() + + socket.onOpen(webSocket, response()) + + verify(webSocket).close(eq(1_000), any()) + verify(connectionCallback, never()).onConnected() + } + + @Test + fun messagesAfterCloseAreDropped() { + val webSocket: WebSocket = mock() + socket.onOpen(webSocket, response()) + socket.closeQuietly() + + socket.onMessage(webSocket, """{"version":2,"method":"reload"}""") + + verify(messageCallback, never()).onMessage(any()) + } + + @Test + fun closeWaitsForARunningMessageHandler() { + val events = CopyOnWriteArrayList() + val handlerStarted = CountDownLatch(1) + val releaseHandler = CountDownLatch(1) + doAnswer { + handlerStarted.countDown() + check(releaseHandler.await(10, TimeUnit.SECONDS)) { "Handler was not released" } + events.add("handler returned") + } + .whenever(messageCallback) + .onMessage(any()) + + val receive = Thread { socket.onMessage(mock(), """{"version":2,"method":"reload"}""") } + val close = Thread { + socket.closeQuietly() + events.add("close returned") + } + receive.start() + try { + assertThat(handlerStarted.await(10, TimeUnit.SECONDS)).isTrue() + close.start() + val deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(10) + while (close.state != Thread.State.BLOCKED) { + check(close.isAlive) { "Close returned while a handler was running" } + check(System.nanoTime() < deadline) { "Close never waited for the handler" } + Thread.yield() + } + } finally { + releaseHandler.countDown() + receive.join(10_000) + close.join(10_000) + } + + assertThat(events).containsExactly("handler returned", "close returned") + } + + private fun response(): Response = + Response.Builder() + .request(Request.Builder().url(URL).build()) + .protocol(Protocol.HTTP_1_1) + .code(101) + .message("Switching Protocols") + .build() + + private companion object { + private const val URL = "ws://127.0.0.1:8081/message" + } +}