diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 0f6f4381f282..9d3921cab7ac 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -89,6 +89,7 @@ import okio.BufferedSource import okio.ByteString import okio.Sink import okio.Timeout +import okio.asOkioSocket import okio.buffer import okio.sink import okio.source @@ -604,9 +605,16 @@ public class MockWebServer : Closeable { "Upgrade".equals(request.headers["Connection"], ignoreCase = true) && "websocket".equals(request.headers["Upgrade"], ignoreCase = true) val responseWantsWebSockets = response.webSocketListener != null + val requestWantsTcp = + "Upgrade".equals(request.headers["Connection"], ignoreCase = true) && + "tcp".equals(request.headers["Upgrade"], ignoreCase = true) + val responseWantsStream = response.socketHandler != null if (requestWantsWebSockets && responseWantsWebSockets) { handleWebSocketUpgrade(socket, source, sink, request, response) reuseSocket = false + } else if (requestWantsTcp && responseWantsStream) { + writeHttpResponse(socket, sink, response) + reuseSocket = false } else { writeHttpResponse(socket, sink, response) } @@ -865,6 +873,11 @@ public class MockWebServer : Closeable { writeHeaders(sink, response.headers) + if (response.socketHandler != null) { + response.socketHandler.handle(socket.asOkioSocket()) + return + } + val body = response.body ?: return socket.sleepWhileOpen(response.bodyDelayNanos) val responseBodySink = diff --git a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt index 956a0e824f59..7dd531c4ee2e 100644 --- a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt +++ b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt @@ -82,7 +82,7 @@ class MockSocketHandler : SocketHandler { @JvmOverloads fun sendResponse( s: String, - responseSent: CountDownLatch = CountDownLatch(0), + responseSent: CountDownLatch = CountDownLatch(1), ) = apply { actions += { stream -> stream.sink.writeUtf8(s) diff --git a/okhttp/api/android/okhttp.api b/okhttp/api/android/okhttp.api index b1ee782b62ac..e0c2d0e41b89 100644 --- a/okhttp/api/android/okhttp.api +++ b/okhttp/api/android/okhttp.api @@ -1120,6 +1120,7 @@ public final class okhttp3/Response : java/io/Closeable { public final fun receivedResponseAtMillis ()J public final fun request ()Lokhttp3/Request; public final fun sentRequestAtMillis ()J + public final fun socket ()Lokio/Socket; public fun toString ()Ljava/lang/String; public final fun trailers ()Lokhttp3/Headers; } @@ -1142,6 +1143,7 @@ public class okhttp3/Response$Builder { public fun removeHeader (Ljava/lang/String;)Lokhttp3/Response$Builder; public fun request (Lokhttp3/Request;)Lokhttp3/Response$Builder; public fun sentRequestAtMillis (J)Lokhttp3/Response$Builder; + public fun socket (Lokio/Socket;)Lokhttp3/Response$Builder; public fun trailers (Lokhttp3/TrailersSource;)Lokhttp3/Response$Builder; } diff --git a/okhttp/api/jvm/okhttp.api b/okhttp/api/jvm/okhttp.api index b1ee782b62ac..e0c2d0e41b89 100644 --- a/okhttp/api/jvm/okhttp.api +++ b/okhttp/api/jvm/okhttp.api @@ -1120,6 +1120,7 @@ public final class okhttp3/Response : java/io/Closeable { public final fun receivedResponseAtMillis ()J public final fun request ()Lokhttp3/Request; public final fun sentRequestAtMillis ()J + public final fun socket ()Lokio/Socket; public fun toString ()Ljava/lang/String; public final fun trailers ()Lokhttp3/Headers; } @@ -1142,6 +1143,7 @@ public class okhttp3/Response$Builder { public fun removeHeader (Ljava/lang/String;)Lokhttp3/Response$Builder; public fun request (Lokhttp3/Request;)Lokhttp3/Response$Builder; public fun sentRequestAtMillis (J)Lokhttp3/Response$Builder; + public fun socket (Lokio/Socket;)Lokhttp3/Response$Builder; public fun trailers (Lokhttp3/TrailersSource;)Lokhttp3/Response$Builder; } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Response.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Response.kt index 427c18cc7af6..3ae3fe8e0d45 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Response.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Response.kt @@ -29,6 +29,7 @@ import okhttp3.internal.http.HTTP_PERM_REDIRECT import okhttp3.internal.http.HTTP_TEMP_REDIRECT import okhttp3.internal.http.parseChallenges import okio.Buffer +import okio.Socket /** * An HTTP response. Instances of this class are not immutable: the response body is a one-shot @@ -77,6 +78,10 @@ class Response internal constructor( * all instances of [ResponseBody]. */ @get:JvmName("body") val body: ResponseBody, + /** + * Non-null if this response is a successful upgrade ... + */ + @get:JvmName("socket") val socket: Socket?, /** * Returns the raw response received from the network. Will be null if this response didn't use * the network, such as when the response is fully cached. The body of the returned response @@ -353,6 +358,7 @@ class Response internal constructor( internal var handshake: Handshake? = null internal var headers: Headers.Builder internal var body: ResponseBody = ResponseBody.EMPTY + internal var socket: Socket? = null internal var networkResponse: Response? = null internal var cacheResponse: Response? = null internal var priorResponse: Response? = null @@ -373,6 +379,7 @@ class Response internal constructor( this.handshake = response.handshake this.headers = response.headers.newBuilder() this.body = response.body + this.socket = response.socket this.networkResponse = response.networkResponse this.cacheResponse = response.cacheResponse this.priorResponse = response.priorResponse @@ -446,6 +453,11 @@ class Response internal constructor( this.body = body } + open fun socket(socket: Socket) = + apply { + this.socket = socket + } + open fun networkResponse(networkResponse: Response?) = apply { checkSupportResponse("networkResponse", networkResponse) @@ -503,6 +515,7 @@ class Response internal constructor( handshake, headers.build(), body, + socket, networkResponse, cacheResponse, priorResponse, diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt index c98e7296b2f9..9c069c7aa5b0 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt @@ -30,6 +30,7 @@ import okio.Buffer import okio.ForwardingSink import okio.ForwardingSource import okio.Sink +import okio.Socket import okio.Source import okio.buffer @@ -157,6 +158,11 @@ class Exchange( ) } + fun newHttpSocket(): Socket { + call.timeoutEarlyExit() + return (codec.carrier as RealConnection).newHttpSocket() + } + fun noNewExchangesOnConnection() { codec.carrier.noNewExchanges() } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt index a048f0cb81e9..ec2fe0c1d27b 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt @@ -58,6 +58,7 @@ import okio.BufferedSource import okio.Sink import okio.Source import okio.Timeout +import okio.asOkioSocket import okio.buffer /** @@ -218,7 +219,7 @@ class RealConnection internal constructor( // 2. The routes must share an IP address. if (routes == null || !routeMatchesAny(routes)) return false - // 3. This connection's server certificate's must cover the new host. + // 3. This connection's server certificates must cover the new host. if (address.hostnameVerifier !== OkHostnameVerifier) return false if (!supportsUrl(address.url)) return false @@ -312,6 +313,12 @@ class RealConnection internal constructor( } } + internal fun newHttpSocket(): okio.Socket { + socket.soTimeout = 0 + noNewExchanges() + return socket.asOkioSocket() + } + override fun route(): Route = route override fun cancel() { diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index 6d0145f61285..72e791fe8d68 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -128,9 +128,21 @@ class CallServerInterceptor( exchange.responseHeadersEnd(response) response = - if (forWebSocket && code == 101) { - // Connection is upgrading, but we need to ensure interceptors see a non-null response body. - response.stripBody() + if (code == HTTP_SWITCHING_PROTOCOLS && + "upgrade".equals(response.request.header("Connection"), ignoreCase = true) && + "upgrade".equals(response.header("Connection"), ignoreCase = true) + ) { + if (forWebSocket) { + // Connection is upgrading, but we need to ensure interceptors see a non-null response body. + response.stripBody() + } else { + // Generic case to return the raw socket. + response + .stripBody() + .newBuilder() + .socket(exchange.newHttpSocket()) + .build() + } } else { val responseBody = exchange.openResponseBody(response) response diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt index 4df249a2dd49..a47546399eb7 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt @@ -225,7 +225,7 @@ fun Response.promisesBody(): Boolean { } val responseCode = code - if ((responseCode < HTTP_CONTINUE || responseCode >= 200) && + if ((responseCode < HTTP_CONTINUE || responseCode == HTTP_SWITCHING_PROTOCOLS || responseCode >= 200) && responseCode != HTTP_NO_CONTENT && responseCode != HTTP_NOT_MODIFIED ) { diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpStatusCodes.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpStatusCodes.kt index d8a3fcce575d..7ffe3d18c238 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpStatusCodes.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpStatusCodes.kt @@ -35,6 +35,9 @@ package okhttp3.internal.http /** `100 Continue` (HTTP/1.1 - RFC 7231) */ const val HTTP_CONTINUE = 100 +/** `101 Switching Protocols` (HTTP/1.1 - RFC 9110) */ +const val HTTP_SWITCHING_PROTOCOLS = 101 + /** `102 Processing` (WebDAV - RFC 2518) */ const val HTTP_PROCESSING = 102 diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt new file mode 100644 index 000000000000..417e22b8d3b6 --- /dev/null +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -0,0 +1,243 @@ +/* + * Copyright (C) 2025 Square, Inc. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package okhttp3.internal.http + +import assertk.assertThat +import assertk.assertions.containsExactly +import assertk.assertions.isEmpty +import assertk.assertions.isEqualTo +import assertk.assertions.isInstanceOf +import assertk.assertions.isNotNull +import java.util.concurrent.BlockingQueue +import java.util.concurrent.LinkedBlockingQueue +import java.util.concurrent.TimeUnit +import mockwebserver3.MockResponse +import mockwebserver3.MockWebServer +import mockwebserver3.junit5.StartStop +import okhttp3.Cache +import okhttp3.Headers.Companion.headersOf +import okhttp3.MediaType.Companion.toMediaType +import okhttp3.OkHttpClient +import okhttp3.OkHttpClientTestRule +import okhttp3.RecordingCallback +import okhttp3.RecordingEventListener +import okhttp3.Request +import okhttp3.TestLogHandler +import okhttp3.internal.UnreadableResponseBody +import okhttp3.internal.closeQuietly +import okhttp3.internal.duplex.MockSocketHandler +import okhttp3.okio.LoggingFilesystem +import okhttp3.testing.PlatformRule +import okio.Path.Companion.toPath +import okio.Socket +import okio.buffer +import okio.fakefilesystem.FakeFileSystem +import okio.use +import org.junit.jupiter.api.AfterEach +import org.junit.jupiter.api.BeforeEach +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.extension.RegisterExtension + +class HttpUpgradesTest { + private val fileSystem = FakeFileSystem() + + @RegisterExtension + val platform = PlatformRule() + + @RegisterExtension + val clientTestRule = OkHttpClientTestRule() + + @RegisterExtension + val testLogHandler = TestLogHandler(OkHttpClient::class.java) + + @StartStop + private val server = MockWebServer() + + @StartStop + private val server2 = MockWebServer() + + private var listener = RecordingEventListener() + private val handshakeCertificates = platform.localhostHandshakeCertificates() + private var client = + clientTestRule + .newClientBuilder() + .eventListenerFactory(clientTestRule.wrap(listener)) + .build() + private val callback = RecordingCallback() + private val cache = + Cache( + fileSystem = LoggingFilesystem(fileSystem), + directory = "/cache".toPath(), + maxSize = Int.MAX_VALUE.toLong(), + ) + + @BeforeEach + fun setUp() { + } + + @AfterEach + @Throws(Exception::class) + fun tearDown() { + } + + @Test + fun upgradeRefusedByServer() { + server.enqueue(MockResponse(body = "normal request")) + val requestWithUpgrade = + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .header("Upgrade", "tcp") + .build() + val response = client.newCall(requestWithUpgrade).execute() + response.body.string() + assertThat(response.code).isEqualTo(200) + } + + @Test + fun upgradesOnReusedConnection() { + server.enqueue(MockResponse(body = "normal request")) + server.enqueue( + MockResponse + .Builder() + .code(HTTP_SWITCHING_PROTOCOLS) + .headers( + headersOf( + "Connection", + "upgrade", + "Upgrade", + "tcp", + "Content-Type", + "text/plain; charset=UTF-8", + ), + ).socketHandler(MockSocketHandler()) + .build(), + ) + val request = Request(server.url("/")) + val requestWithUpgrade = + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .header("Upgrade", "tcp") + .build() + assertConnectionReused(request, requestWithUpgrade) + } + + // copied from okhttp3.ConnectionReuseTest.assertConnectionReused + private fun assertConnectionReused(vararg requests: Request?) { + for (i in requests.indices) { + val response = client.newCall(requests[i]!!).execute() + if (response.code == HTTP_SWITCHING_PROTOCOLS) { + response.exchange!!.cancel() + } else { + response.body.string() // Discard the response body. + } + assertThat(server.takeRequest().exchangeIndex).isEqualTo(i) + } + } + + @Test + fun upgradeConnection() { + val mockStreamHandler = + MockSocketHandler() + .receiveRequest("request A\n") + .sendResponse("response B\n") + .receiveRequest("request C\n") + .sendResponse("response D\n") + .sendResponse("response E\n") + .receiveRequest("response F\n") + .exhaustRequest() + .exhaustResponse() + server.enqueue( + MockResponse + .Builder() + .code(HTTP_SWITCHING_PROTOCOLS) + .headers( + headersOf( + "Connection", + "upgrade", + "Upgrade", + "tcp", + "Content-Type", + "text/plain; charset=UTF-8", +// "Content-Type", "application/vnd.docker.raw-stream", + ), + ).socketHandler(mockStreamHandler) + .build(), + ) + val call = + client.newCall( + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .header("Upgrade", "tcp") + // .post(...) + .build(), + ) + + var socket: Socket? + val received: BlockingQueue = LinkedBlockingQueue() + + call.execute().use { response -> + assertThat(response.code).isEqualTo(HTTP_SWITCHING_PROTOCOLS) + assertThat(response.headers("Connection").first()).isEqualTo("upgrade", true) + assertThat(response.headers("Upgrade").first()).isEqualTo("tcp", true) + assertThat(response.headers("Content-Type").first().toMediaType()).isEqualTo("text/plain; charset=UTF-8".toMediaType()) +// assertThat(response.headers("Content-Type").first().toMediaType()).isEqualTo("application/vnd.docker.raw-stream".toMediaType()) + assertThat(response.headers("Content-Length")).isEmpty() + assertThat(response.body).isInstanceOf() + + socket = response.socket + assertThat(socket).isNotNull() + + val reader = socket!!.source + val readerThread = + object : Thread("reader") { + override fun run() { + try { + var line: String? + while (reader.buffer().readUtf8Line().also { line = it } != null) { + received.add(line) + } + } catch (e: Exception) { + reader.closeQuietly() + } + } + } + readerThread.start() + + val writer = socket.sink + writer.buffer().writeUtf8("request A\n").flush() + writer.buffer().writeUtf8("request C\n").flush() + writer.buffer().writeUtf8("response F\n").flush() + } + + val responses = mutableListOf() + responses.add(received.poll(2, TimeUnit.SECONDS)) + responses.add(received.poll(2, TimeUnit.SECONDS)) + responses.add(received.poll(2, TimeUnit.SECONDS)) + assertThat(responses).containsExactly( + "response B", + "response D", + "response E", + ) + + socket?.cancel() + } +}