From 889b147f8b401a70f0e139ccf6778ccbead3c4f1 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Tue, 10 Jun 2025 20:43:42 +0200 Subject: [PATCH 01/10] Give the CountDownLatch something to count --- .../main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) 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) From 9009b059ad36411ab219ae82942ad4d583932649 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Tue, 10 Jun 2025 20:44:53 +0200 Subject: [PATCH 02/10] Fix typo --- .../kotlin/okhttp3/internal/connection/RealConnection.kt | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt index a048f0cb81e9..e0cf6e318220 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt @@ -218,7 +218,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 From 79e4373f8169f2893daa6d3792faca265c2eb9f7 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Wed, 7 Feb 2024 23:18:40 +0100 Subject: [PATCH 03/10] Connection Upgrade for HTTP1 status code 101 See https://datatracker.ietf.org/doc/html/rfc9110#name-101-switching-protocols --- .../kotlin/mockwebserver3/MockWebServer.kt | 24 +++++ .../kotlin/okhttp3/Response.kt | 13 +++ .../okhttp3/internal/connection/Exchange.kt | 6 ++ .../internal/connection/RealConnection.kt | 18 ++++ .../internal/http/CallServerInterceptor.kt | 45 ++++++--- .../okhttp3/internal/http/HttpHeaders.kt | 2 +- .../okhttp3/internal/http/HttpStatusCodes.kt | 3 + okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt | 97 +++++++++++++++++++ 8 files changed, 191 insertions(+), 17 deletions(-) diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 0f6f4381f282..9bbfe1cff736 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -88,6 +88,7 @@ import okio.BufferedSink import okio.BufferedSource import okio.ByteString import okio.Sink +import okio.Source import okio.Timeout import okio.buffer import okio.sink @@ -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,22 @@ public class MockWebServer : Closeable { writeHeaders(sink, response.headers) + if (response.socketHandler != null) { + response.socketHandler.handle( + object : okio.Socket { + override val source: Source + get() = socket.source() + override val sink: Sink + get() = socket.sink() + + override fun cancel() { + socket.closeQuietly() + } + }, + ) + return + } + val body = response.body ?: return socket.sleepWhileOpen(response.bodyDelayNanos) val responseBodySink = 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..43e5f7e7c0b3 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 newHttpStreams(): Socket { + call.timeoutEarlyExit() + return (codec.carrier as RealConnection).newHttpSocket(this) + } + 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 e0cf6e318220..cfa6d446316e 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt @@ -59,6 +59,8 @@ import okio.Sink import okio.Source import okio.Timeout import okio.buffer +import okio.sink +import okio.source /** * A connection to a remote web server capable of carrying 1 or more concurrent streams. @@ -312,6 +314,22 @@ class RealConnection internal constructor( } } + internal fun newHttpSocket(exchange: Exchange): okio.Socket { + socket.soTimeout = 0 + noNewExchanges() + return object : okio.Socket { + override val source: Source + get() = socket.source() + + override val sink: Sink + get() = socket.sink() + + override fun cancel() { + exchange.cancel() + } + } + } + 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..479eec134816 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -128,27 +128,40 @@ class CallServerInterceptor( exchange.responseHeadersEnd(response) response = - if (forWebSocket && code == 101) { + if (forWebSocket && code == HTTP_SWITCHING_PROTOCOLS) { // Connection is upgrading, but we need to ensure interceptors see a non-null response body. response.stripBody() } else { - val responseBody = exchange.openResponseBody(response) - response - .newBuilder() - .body(responseBody) - .trailers( - object : TrailersSource { - override fun peek() = exchange.peekTrailers() + if (code == HTTP_SWITCHING_PROTOCOLS && + "upgrade".equals(response.request.header("Connection"), ignoreCase = true) && + "upgrade".equals(response.header("Connection"), ignoreCase = true) && + "tcp".equals(response.request.header("Upgrade"), ignoreCase = true) && + "tcp".equals(response.header("Upgrade"), ignoreCase = true) + ) { + response + .stripBody() + .newBuilder() + .socket(exchange.newHttpStreams()) + .build() + } else { + val responseBody = exchange.openResponseBody(response) + response + .newBuilder() + .body(responseBody) + .trailers( + object : TrailersSource { + override fun peek() = exchange.peekTrailers() - override fun get(): Headers { - val source = responseBody.source() - if (source.isOpen) { - source.skipAll() + override fun get(): Headers { + val source = responseBody.source() + if (source.isOpen) { + source.skipAll() + } + return peek() ?: error("null trailers after exhausting response body?!") } - return peek() ?: error("null trailers after exhausting response body?!") - } - }, - ).build() + }, + ).build() + } } if ("close".equals(response.request.header("Connection"), ignoreCase = true) || "close".equals(response.header("Connection"), ignoreCase = true) 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/CallTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt index c5db8fd55d3a..9b60c5c613f6 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt @@ -24,9 +24,11 @@ import assertk.assertions.hasMessage import assertk.assertions.hasSize import assertk.assertions.index import assertk.assertions.isCloseTo +import assertk.assertions.isEmpty import assertk.assertions.isEqualTo import assertk.assertions.isFalse import assertk.assertions.isIn +import assertk.assertions.isInstanceOf import assertk.assertions.isLessThan import assertk.assertions.isNotEmpty import assertk.assertions.isNotNull @@ -55,6 +57,7 @@ import java.util.Arrays import java.util.concurrent.BlockingQueue import java.util.concurrent.CountDownLatch import java.util.concurrent.Executors +import java.util.concurrent.LinkedBlockingQueue import java.util.concurrent.SynchronousQueue import java.util.concurrent.TimeUnit import java.util.concurrent.atomic.AtomicBoolean @@ -93,10 +96,13 @@ import okhttp3.TestUtil.awaitGarbageCollection import okhttp3.internal.DoubleInetAddressDns import okhttp3.internal.RecordingOkAuthenticator import okhttp3.internal.USER_AGENT +import okhttp3.internal.UnreadableResponseBody import okhttp3.internal.addHeaderLenient import okhttp3.internal.closeQuietly +import okhttp3.internal.duplex.MockSocketHandler import okhttp3.internal.http.HTTP_EARLY_HINTS import okhttp3.internal.http.HTTP_PROCESSING +import okhttp3.internal.http.HTTP_SWITCHING_PROTOCOLS import okhttp3.internal.http.RecordingProxySelector import okhttp3.java.net.cookiejar.JavaNetCookieJar import okhttp3.okio.LoggingFilesystem @@ -110,6 +116,7 @@ import okio.ByteString import okio.ForwardingSource import okio.GzipSink import okio.Path.Companion.toPath +import okio.Socket import okio.buffer import okio.fakefilesystem.FakeFileSystem import okio.use @@ -4841,6 +4848,96 @@ open class CallTest { } } + @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() + } + private fun makeFailingCall() { val requestBody: RequestBody = object : RequestBody() { From 6e337e3311784d335b719c6d8a639f4a6dab53bf Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Fri, 25 Jul 2025 21:08:28 +0200 Subject: [PATCH 04/10] use socket.asOkioSocket() --- .../main/kotlin/mockwebserver3/MockWebServer.kt | 15 ++------------- .../okhttp3/internal/connection/Exchange.kt | 2 +- .../internal/connection/RealConnection.kt | 17 +++-------------- 3 files changed, 6 insertions(+), 28 deletions(-) diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 9bbfe1cff736..9d3921cab7ac 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -88,8 +88,8 @@ import okio.BufferedSink import okio.BufferedSource import okio.ByteString import okio.Sink -import okio.Source import okio.Timeout +import okio.asOkioSocket import okio.buffer import okio.sink import okio.source @@ -874,18 +874,7 @@ public class MockWebServer : Closeable { writeHeaders(sink, response.headers) if (response.socketHandler != null) { - response.socketHandler.handle( - object : okio.Socket { - override val source: Source - get() = socket.source() - override val sink: Sink - get() = socket.sink() - - override fun cancel() { - socket.closeQuietly() - } - }, - ) + response.socketHandler.handle(socket.asOkioSocket()) return } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt index 43e5f7e7c0b3..5764229551f7 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt @@ -160,7 +160,7 @@ class Exchange( fun newHttpStreams(): Socket { call.timeoutEarlyExit() - return (codec.carrier as RealConnection).newHttpSocket(this) + return (codec.carrier as RealConnection).newHttpSocket() } fun noNewExchangesOnConnection() { diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt index cfa6d446316e..ec2fe0c1d27b 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt @@ -58,9 +58,8 @@ import okio.BufferedSource import okio.Sink import okio.Source import okio.Timeout +import okio.asOkioSocket import okio.buffer -import okio.sink -import okio.source /** * A connection to a remote web server capable of carrying 1 or more concurrent streams. @@ -314,20 +313,10 @@ class RealConnection internal constructor( } } - internal fun newHttpSocket(exchange: Exchange): okio.Socket { + internal fun newHttpSocket(): okio.Socket { socket.soTimeout = 0 noNewExchanges() - return object : okio.Socket { - override val source: Source - get() = socket.source() - - override val sink: Sink - get() = socket.sink() - - override fun cancel() { - exchange.cancel() - } - } + return socket.asOkioSocket() } override fun route(): Route = route From 2d4d6784f8b28047753144d725657bad9d7b9c64 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Fri, 25 Jul 2025 21:20:13 +0200 Subject: [PATCH 05/10] add socket to api --- okhttp/api/android/okhttp.api | 2 ++ okhttp/api/jvm/okhttp.api | 2 ++ 2 files changed, 4 insertions(+) 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; } From ffce1d1556c8f183683070210524d8aa8523c7fe Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Fri, 25 Jul 2025 21:39:51 +0200 Subject: [PATCH 06/10] remove the specific check on tcp connection upgrades --- .../internal/http/CallServerInterceptor.kt | 53 +++++++++---------- 1 file changed, 26 insertions(+), 27 deletions(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index 479eec134816..ad7941c65a4e 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -128,40 +128,39 @@ class CallServerInterceptor( exchange.responseHeadersEnd(response) response = - if (forWebSocket && code == HTTP_SWITCHING_PROTOCOLS) { - // Connection is upgrading, but we need to ensure interceptors see a non-null response body. - response.stripBody() - } else { - if (code == HTTP_SWITCHING_PROTOCOLS && - "upgrade".equals(response.request.header("Connection"), ignoreCase = true) && - "upgrade".equals(response.header("Connection"), ignoreCase = true) && - "tcp".equals(response.request.header("Upgrade"), ignoreCase = true) && - "tcp".equals(response.header("Upgrade"), ignoreCase = true) - ) { + 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.newHttpStreams()) .build() - } else { - val responseBody = exchange.openResponseBody(response) - response - .newBuilder() - .body(responseBody) - .trailers( - object : TrailersSource { - override fun peek() = exchange.peekTrailers() + } + } else { + val responseBody = exchange.openResponseBody(response) + response + .newBuilder() + .body(responseBody) + .trailers( + object : TrailersSource { + override fun peek() = exchange.peekTrailers() - override fun get(): Headers { - val source = responseBody.source() - if (source.isOpen) { - source.skipAll() - } - return peek() ?: error("null trailers after exhausting response body?!") + override fun get(): Headers { + val source = responseBody.source() + if (source.isOpen) { + source.skipAll() } - }, - ).build() - } + return peek() ?: error("null trailers after exhausting response body?!") + } + }, + ).build() } if ("close".equals(response.request.header("Connection"), ignoreCase = true) || "close".equals(response.header("Connection"), ignoreCase = true) From 335d69ef417ca107c3ed9a19d2e6b1fbed46cba3 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Fri, 25 Jul 2025 22:05:20 +0200 Subject: [PATCH 07/10] create a dedicated HttpUpgradesTest --- okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt | 97 --------- .../okhttp3/internal/http/HttpUpgradesTest.kt | 185 ++++++++++++++++++ 2 files changed, 185 insertions(+), 97 deletions(-) create mode 100644 okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt index 9b60c5c613f6..c5db8fd55d3a 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/CallTest.kt @@ -24,11 +24,9 @@ import assertk.assertions.hasMessage import assertk.assertions.hasSize import assertk.assertions.index import assertk.assertions.isCloseTo -import assertk.assertions.isEmpty import assertk.assertions.isEqualTo import assertk.assertions.isFalse import assertk.assertions.isIn -import assertk.assertions.isInstanceOf import assertk.assertions.isLessThan import assertk.assertions.isNotEmpty import assertk.assertions.isNotNull @@ -57,7 +55,6 @@ import java.util.Arrays import java.util.concurrent.BlockingQueue import java.util.concurrent.CountDownLatch import java.util.concurrent.Executors -import java.util.concurrent.LinkedBlockingQueue import java.util.concurrent.SynchronousQueue import java.util.concurrent.TimeUnit import java.util.concurrent.atomic.AtomicBoolean @@ -96,13 +93,10 @@ import okhttp3.TestUtil.awaitGarbageCollection import okhttp3.internal.DoubleInetAddressDns import okhttp3.internal.RecordingOkAuthenticator import okhttp3.internal.USER_AGENT -import okhttp3.internal.UnreadableResponseBody import okhttp3.internal.addHeaderLenient import okhttp3.internal.closeQuietly -import okhttp3.internal.duplex.MockSocketHandler import okhttp3.internal.http.HTTP_EARLY_HINTS import okhttp3.internal.http.HTTP_PROCESSING -import okhttp3.internal.http.HTTP_SWITCHING_PROTOCOLS import okhttp3.internal.http.RecordingProxySelector import okhttp3.java.net.cookiejar.JavaNetCookieJar import okhttp3.okio.LoggingFilesystem @@ -116,7 +110,6 @@ import okio.ByteString import okio.ForwardingSource import okio.GzipSink import okio.Path.Companion.toPath -import okio.Socket import okio.buffer import okio.fakefilesystem.FakeFileSystem import okio.use @@ -4848,96 +4841,6 @@ open class CallTest { } } - @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() - } - private fun makeFailingCall() { val requestBody: RequestBody = object : RequestBody() { 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..04a1c2969eb3 --- /dev/null +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -0,0 +1,185 @@ +/* + * 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 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() + } +} From dd09d8d837f5f4ef7c4a88aab5a9be0cf972a645 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Fri, 25 Jul 2025 22:14:06 +0200 Subject: [PATCH 08/10] chore --- .../kotlin/okhttp3/internal/connection/Exchange.kt | 2 +- .../kotlin/okhttp3/internal/http/CallServerInterceptor.kt | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt index 5764229551f7..9c069c7aa5b0 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt @@ -158,7 +158,7 @@ class Exchange( ) } - fun newHttpStreams(): Socket { + fun newHttpSocket(): Socket { call.timeoutEarlyExit() return (codec.carrier as RealConnection).newHttpSocket() } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index ad7941c65a4e..72e791fe8d68 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -140,7 +140,7 @@ class CallServerInterceptor( response .stripBody() .newBuilder() - .socket(exchange.newHttpStreams()) + .socket(exchange.newHttpSocket()) .build() } } else { From f2daee6d396c7dd82912bd33852927ac3350d18f Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Sun, 27 Jul 2025 14:49:10 +0200 Subject: [PATCH 09/10] add more test cases --- .../okhttp3/internal/http/HttpUpgradesTest.kt | 55 +++++++++++++++++++ 1 file changed, 55 insertions(+) diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 04a1c2969eb3..0a61102d7b16 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -93,6 +93,61 @@ class HttpUpgradesTest { 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 = From ab7c625c3825c07545c346ac8ca82edaae06e8df Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Sun, 27 Jul 2025 15:00:16 +0200 Subject: [PATCH 10/10] spotlessApply --- .../okhttp3/internal/http/HttpUpgradesTest.kt | 31 ++++++++++--------- 1 file changed, 17 insertions(+), 14 deletions(-) diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 0a61102d7b16..417e22b8d3b6 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -96,12 +96,13 @@ class HttpUpgradesTest { @Test fun upgradeRefusedByServer() { server.enqueue(MockResponse(body = "normal request")) - val requestWithUpgrade = Request - .Builder() - .url(server.url("/")) - .header("Connection", "upgrade") - .header("Upgrade", "tcp") - .build() + 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) @@ -123,15 +124,17 @@ class HttpUpgradesTest { "Content-Type", "text/plain; charset=UTF-8", ), - ).socketHandler( MockSocketHandler()) - .build()) + ).socketHandler(MockSocketHandler()) + .build(), + ) val request = Request(server.url("/")) - val requestWithUpgrade = Request - .Builder() - .url(server.url("/")) - .header("Connection", "upgrade") - .header("Upgrade", "tcp") - .build() + val requestWithUpgrade = + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .header("Upgrade", "tcp") + .build() assertConnectionReused(request, requestWithUpgrade) }