From 889b147f8b401a70f0e139ccf6778ccbead3c4f1 Mon Sep 17 00:00:00 2001 From: Tobias Gesellchen Date: Tue, 10 Jun 2025 20:43:42 +0200 Subject: [PATCH 01/17] 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/17] 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/17] 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/17] 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/17] 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/17] 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/17] 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/17] 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/17] 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/17] 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) } From 0756bf9d13ec05a8f63a69508ad3b5c4241a0c00 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Sun, 27 Jul 2025 12:11:58 -0400 Subject: [PATCH 11/17] Use more MockWebServer features in tests --- .../okhttp3/internal/http/HttpUpgradesTest.kt | 251 +++++++----------- 1 file changed, 92 insertions(+), 159 deletions(-) diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 417e22b8d3b6..029c77d707e3 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -16,59 +16,37 @@ 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 assertk.assertions.isNull +import assertk.assertions.isTrue +import kotlin.test.assertFailsWith 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.Protocol import okhttp3.RecordingEventListener +import okhttp3.RecordingHostnameVerifier 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.ProtocolException 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 = @@ -76,21 +54,48 @@ class HttpUpgradesTest { .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() { + @Test + fun upgrade() { + val socketHandler = MockSocketHandler() + .apply { + receiveRequest("client says hello\n") + sendResponse("server says hello\n") + receiveRequest("client says goodbye\n") + sendResponse("server says goodbye\n") + exhaustResponse() + exhaustRequest() + } + server.enqueue(socketHandler.upgradeResponse()) + + client.newCall( + upgradeRequest() + ).execute().use { response -> + assertThat(response.code).isEqualTo(HTTP_SWITCHING_PROTOCOLS) + val socket = response.socket!! + socket.sink.buffer().use { sink -> + socket.source.buffer().use { source -> + sink.writeUtf8("client says hello\n") + sink.flush() + + assertThat(source.readUtf8Line()).isEqualTo("server says hello") + + sink.writeUtf8("client says goodbye\n") + sink.flush() + + assertThat(source.readUtf8Line()).isEqualTo("server says goodbye") + + assertThat(source.exhausted()).isTrue() + } + } + socketHandler.awaitSuccess() + } } - @AfterEach - @Throws(Exception::class) - fun tearDown() { + @Test + fun upgradeHttps() { + enableTls(Protocol.HTTP_1_1) + upgrade() } @Test @@ -103,31 +108,18 @@ class HttpUpgradesTest { .header("Connection", "upgrade") .header("Upgrade", "tcp") .build() - val response = client.newCall(requestWithUpgrade).execute() - response.body.string() - assertThat(response.code).isEqualTo(200) + client.newCall(requestWithUpgrade).execute().use { response -> + assertThat(response.code).isEqualTo(200) + assertThat(response.socket).isNull() + assertThat(response.body.string()).isEqualTo("normal request") + } } @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("/")) + fun upgradeForbiddenOnHttp2() { + enableTls(Protocol.HTTP_2, Protocol.HTTP_1_1) + val socketHandler = MockSocketHandler() + server.enqueue(socketHandler.upgradeResponse()) val requestWithUpgrade = Request .Builder() @@ -135,109 +127,50 @@ class HttpUpgradesTest { .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) + assertFailsWith { + client.newCall(requestWithUpgrade).execute() } } @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() + fun upgradesOnReusedConnection() { + server.enqueue(MockResponse(body = "normal request")) + client.newCall(Request(server.url("/"))).execute().use { response -> + assertThat(response.body.string()).isEqualTo("normal request") } - 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", - ) + upgrade() - socket?.cancel() + assertThat(server.takeRequest().connectionIndex).isEqualTo(0) + assertThat(server.takeRequest().connectionIndex).isEqualTo(0) } + + private fun enableTls(vararg protocols: Protocol) { + client = + client + .newBuilder() + .protocols(protocols.toList()) + .sslSocketFactory( + handshakeCertificates.sslSocketFactory(), + handshakeCertificates.trustManager, + ).hostnameVerifier(RecordingHostnameVerifier()) + .build() + server.useHttps(handshakeCertificates.sslSocketFactory()) + server.protocols = protocols.toList() + } + + private fun upgradeRequest() = Request( + url = server.url("/"), + headers = headersOf( + "Connection", "upgrade", + "Upgrade", "tcp", + ) + ) + + private fun MockSocketHandler.upgradeResponse() = MockResponse.Builder() + .code(HTTP_SWITCHING_PROTOCOLS) + .addHeader("Connection", "upgrade") + .addHeader("Upgrade", "tcp") + .socketHandler(this) + .build() } From 605aa93bb77f732323997276389c928e9a17d53e Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Sun, 27 Jul 2025 12:35:54 -0400 Subject: [PATCH 12/17] Forbid HTTP/2 --- .../kotlin/okhttp3/internal/http/CallServerInterceptor.kt | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index 72e791fe8d68..13843d59046c 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -127,8 +127,13 @@ class CallServerInterceptor( exchange.responseHeadersEnd(response) + val isUpgradeCode = code == HTTP_SWITCHING_PROTOCOLS + if (isUpgradeCode && exchange.connection.isMultiplexed) { + throw ProtocolException("Unexpected $HTTP_SWITCHING_PROTOCOLS code on HTTP/2 connection") + } + response = - if (code == HTTP_SWITCHING_PROTOCOLS && + if (isUpgradeCode && "upgrade".equals(response.request.header("Connection"), ignoreCase = true) && "upgrade".equals(response.header("Connection"), ignoreCase = true) ) { From 6d73dd0827041cb53b3bf6a755b56fea9aabde30 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Sun, 27 Jul 2025 14:14:48 -0400 Subject: [PATCH 13/17] Hook into Exchange for connection pool management --- .../okhttp3/internal/connection/Exchange.kt | 15 +++- .../internal/connection/RealConnection.kt | 7 +- .../internal/http/CallServerInterceptor.kt | 2 +- .../okhttp3/internal/http/ExchangeCodec.kt | 6 ++ .../internal/http1/Http1ExchangeCodec.kt | 7 ++ .../internal/http2/Http2ExchangeCodec.kt | 6 ++ .../okhttp3/internal/http/HttpUpgradesTest.kt | 85 ++++++++++--------- 7 files changed, 82 insertions(+), 46 deletions(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt index 9c069c7aa5b0..89d30f22dc08 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/Exchange.kt @@ -158,9 +158,20 @@ class Exchange( ) } - fun newHttpSocket(): Socket { + fun upgradeToSocket(): Socket { call.timeoutEarlyExit() - return (codec.carrier as RealConnection).newHttpSocket() + (codec.carrier as RealConnection).useAsSocket() + + eventListener.requestBodyStart(call) + + return object : Socket { + override fun cancel() { + this@Exchange.cancel() + } + + override val sink = RequestBodySink(codec.socketSink, -1L) + override val source = ResponseBodySource(codec.socketSource, -1L) + } } 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 ec2fe0c1d27b..f00dda51043c 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt @@ -58,7 +58,6 @@ import okio.BufferedSource import okio.Sink import okio.Source import okio.Timeout -import okio.asOkioSocket import okio.buffer /** @@ -295,8 +294,7 @@ class RealConnection internal constructor( @Throws(SocketException::class) internal fun newWebSocketStreams(exchange: Exchange): RealWebSocket.Streams { - socket.soTimeout = 0 - noNewExchanges() + useAsSocket() return object : RealWebSocket.Streams(true, source, sink) { override fun close() { exchange.bodyComplete( @@ -313,10 +311,9 @@ class RealConnection internal constructor( } } - internal fun newHttpSocket(): okio.Socket { + internal fun useAsSocket() { socket.soTimeout = 0 noNewExchanges() - return socket.asOkioSocket() } override fun route(): Route = route diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index 13843d59046c..df9b8c47619e 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -145,7 +145,7 @@ class CallServerInterceptor( response .stripBody() .newBuilder() - .socket(exchange.newHttpSocket()) + .socket(exchange.upgradeToSocket()) .build() } } else { diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/ExchangeCodec.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/ExchangeCodec.kt index 826093569aad..f1edb85659f7 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/ExchangeCodec.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/ExchangeCodec.kt @@ -32,6 +32,12 @@ interface ExchangeCodec { /** Returns true if the response body and (possibly empty) trailers have been received. */ val isResponseComplete: Boolean + /** The source when this is the subject of a protocol upgrade or CONNECT. */ + val socketSink: Sink + + /** The sink when this is the subject of a protocol upgrade or CONNECT. */ + val socketSource: Source + /** Returns an output stream where the request body can be streamed. */ @Throws(IOException::class) fun createRequestBody( diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http1/Http1ExchangeCodec.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http1/Http1ExchangeCodec.kt index 0826e396592e..bddeac9a503d 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http1/Http1ExchangeCodec.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http1/Http1ExchangeCodec.kt @@ -34,6 +34,7 @@ import okhttp3.internal.http.RequestLine import okhttp3.internal.http.StatusLine import okhttp3.internal.http.promisesBody import okhttp3.internal.http.receiveHeaders +import okhttp3.internal.http1.Http1ExchangeCodec.Companion.TRAILERS_RESPONSE_BODY_TRUNCATED import okhttp3.internal.skipAll import okio.Buffer import okio.BufferedSink @@ -91,6 +92,12 @@ class Http1ExchangeCodec( override val isResponseComplete: Boolean get() = state == STATE_CLOSED + override val socketSink + get() = sink + + override val socketSource + get() = source + override fun createRequestBody( request: Request, contentLength: Long, diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http2/Http2ExchangeCodec.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http2/Http2ExchangeCodec.kt index e1d0f53d6f80..b3a217df89d0 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http2/Http2ExchangeCodec.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http2/Http2ExchangeCodec.kt @@ -67,6 +67,12 @@ class Http2ExchangeCodec( override val isResponseComplete: Boolean get() = stream?.isSourceComplete == true + override val socketSink + get() = stream!!.sink + + override val socketSource + get() = stream!!.source + override fun createRequestBody( request: Request, contentLength: Long, diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 029c77d707e3..225eebc0048e 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -57,39 +57,42 @@ class HttpUpgradesTest { @Test fun upgrade() { - val socketHandler = MockSocketHandler() - .apply { - receiveRequest("client says hello\n") - sendResponse("server says hello\n") - receiveRequest("client says goodbye\n") - sendResponse("server says goodbye\n") - exhaustResponse() - exhaustRequest() - } + val socketHandler = + MockSocketHandler() + .apply { + receiveRequest("client says hello\n") + sendResponse("server says hello\n") + receiveRequest("client says goodbye\n") + sendResponse("server says goodbye\n") + exhaustResponse() + exhaustRequest() + } server.enqueue(socketHandler.upgradeResponse()) - client.newCall( - upgradeRequest() - ).execute().use { response -> - assertThat(response.code).isEqualTo(HTTP_SWITCHING_PROTOCOLS) - val socket = response.socket!! - socket.sink.buffer().use { sink -> - socket.source.buffer().use { source -> - sink.writeUtf8("client says hello\n") - sink.flush() + client + .newCall( + upgradeRequest(), + ).execute() + .use { response -> + assertThat(response.code).isEqualTo(HTTP_SWITCHING_PROTOCOLS) + val socket = response.socket!! + socket.sink.buffer().use { sink -> + socket.source.buffer().use { source -> + sink.writeUtf8("client says hello\n") + sink.flush() - assertThat(source.readUtf8Line()).isEqualTo("server says hello") + assertThat(source.readUtf8Line()).isEqualTo("server says hello") - sink.writeUtf8("client says goodbye\n") - sink.flush() + sink.writeUtf8("client says goodbye\n") + sink.flush() - assertThat(source.readUtf8Line()).isEqualTo("server says goodbye") + assertThat(source.readUtf8Line()).isEqualTo("server says goodbye") - assertThat(source.exhausted()).isTrue() + assertThat(source.exhausted()).isTrue() + } } + socketHandler.awaitSuccess() } - socketHandler.awaitSuccess() - } } @Test @@ -159,18 +162,24 @@ class HttpUpgradesTest { server.protocols = protocols.toList() } - private fun upgradeRequest() = Request( - url = server.url("/"), - headers = headersOf( - "Connection", "upgrade", - "Upgrade", "tcp", + private fun upgradeRequest() = + Request( + url = server.url("/"), + headers = + headersOf( + "Connection", + "upgrade", + "Upgrade", + "tcp", + ), ) - ) - - private fun MockSocketHandler.upgradeResponse() = MockResponse.Builder() - .code(HTTP_SWITCHING_PROTOCOLS) - .addHeader("Connection", "upgrade") - .addHeader("Upgrade", "tcp") - .socketHandler(this) - .build() + + private fun MockSocketHandler.upgradeResponse() = + MockResponse + .Builder() + .code(HTTP_SWITCHING_PROTOCOLS) + .addHeader("Connection", "upgrade") + .addHeader("Upgrade", "tcp") + .socketHandler(this) + .build() } From 9cb812c0adf777246d2eb984d2732f25ecbaa6cd Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Sun, 27 Jul 2025 14:30:20 -0400 Subject: [PATCH 14/17] Test events on successful upgrade --- .../internal/http/CallServerInterceptor.kt | 10 +++---- .../okhttp3/internal/http/HttpUpgradesTest.kt | 28 +++++++++++++++++++ 2 files changed, 33 insertions(+), 5 deletions(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index df9b8c47619e..ac0f6d5559c5 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -42,6 +42,7 @@ class CallServerInterceptor( var invokeStartEvent = true var responseBuilder: Response.Builder? = null var sendRequestException: IOException? = null + val isUpgradeRequest = "upgrade".equals(request.header("Connection"), ignoreCase = true) try { exchange.writeRequestHeaders(request) @@ -76,7 +77,7 @@ class CallServerInterceptor( exchange.noNewExchangesOnConnection() } } - } else { + } else if (!isUpgradeRequest) { exchange.noRequestBody() } @@ -132,11 +133,10 @@ class CallServerInterceptor( throw ProtocolException("Unexpected $HTTP_SWITCHING_PROTOCOLS code on HTTP/2 connection") } + val isUpgradeResponse = "upgrade".equals(response.header("Connection"), ignoreCase = true) + // TODO(jwilson): maybe call exchange.noRequestBody() if the upgrade failed? response = - if (isUpgradeCode && - "upgrade".equals(response.request.header("Connection"), ignoreCase = true) && - "upgrade".equals(response.header("Connection"), ignoreCase = true) - ) { + if (isUpgradeCode && isUpgradeRequest && isUpgradeResponse) { if (forWebSocket) { // Connection is upgrading, but we need to ensure interceptors see a non-null response body. response.stripBody() diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 225eebc0048e..e3efab64e3cb 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -16,6 +16,7 @@ package okhttp3.internal.http import assertk.assertThat +import assertk.assertions.containsExactly import assertk.assertions.isEqualTo import assertk.assertions.isNull import assertk.assertions.isTrue @@ -148,6 +149,33 @@ class HttpUpgradesTest { assertThat(server.takeRequest().connectionIndex).isEqualTo(0) } + @Test + fun upgradeEvents() { + upgrade() + + assertThat(listener.recordedEventTypes()).containsExactly( + "CallStart", + "ProxySelectStart", + "ProxySelectEnd", + "DnsStart", + "DnsEnd", + "ConnectStart", + "ConnectEnd", + "ConnectionAcquired", + "RequestHeadersStart", + "RequestHeadersEnd", + "ResponseHeadersStart", + "ResponseHeadersEnd", + "RequestBodyStart", + "FollowUpDecision", + "ResponseBodyStart", + "ResponseBodyEnd", + "RequestBodyEnd", + "ConnectionReleased", + "CallEnd", + ) + } + private fun enableTls(vararg protocols: Protocol) { client = client From 93c2ec3d6ca9afb10c903f5052cc7bd050419fb5 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Mon, 28 Jul 2025 10:10:05 -0400 Subject: [PATCH 15/17] Release the request body after an upgrade fails --- .../kotlin/okhttp3/Request.kt | 7 +++ .../internal/http/CallServerInterceptor.kt | 50 ++++++++++++------- .../okhttp3/internal/http/HttpUpgradesTest.kt | 19 +++++-- 3 files changed, 53 insertions(+), 23 deletions(-) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Request.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Request.kt index fd7a6568ff8d..0f027aee19a9 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Request.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/Request.kt @@ -80,6 +80,13 @@ class Request internal constructor( ), ) + init { + val connectionHeader = headers["Connection"] + require(body == null || !"upgrade".equals(connectionHeader, ignoreCase = true)) { + "expected a null request body with 'Connection: upgrade'" + } + } + fun header(name: String): String? = headers[name] fun headers(name: String): List = headers.values(name) diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index ac0f6d5559c5..7f017edbee30 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -21,10 +21,10 @@ import okhttp3.Headers import okhttp3.Interceptor import okhttp3.Response import okhttp3.TrailersSource +import okhttp3.internal.UnreadableResponseBody import okhttp3.internal.connection.Exchange import okhttp3.internal.http2.ConnectionShutdownException import okhttp3.internal.skipAll -import okhttp3.internal.stripBody import okio.buffer /** This is the last interceptor in the chain. It makes a network call to the server. */ @@ -42,11 +42,13 @@ class CallServerInterceptor( var invokeStartEvent = true var responseBuilder: Response.Builder? = null var sendRequestException: IOException? = null - val isUpgradeRequest = "upgrade".equals(request.header("Connection"), ignoreCase = true) + val hasRequestBody = HttpMethod.permitsRequestBody(request.method) && requestBody != null + val isUpgradeRequest = !hasRequestBody + && "upgrade".equals(request.header("Connection"), ignoreCase = true) try { exchange.writeRequestHeaders(request) - if (HttpMethod.permitsRequestBody(request.method) && requestBody != null) { + if (hasRequestBody) { // If there's a "Expect: 100-continue" header on the request, wait for a "HTTP/1.1 100 // Continue" response before transmitting the request body. If we don't get that, return // what we did get (such as a 4xx response) without ever transmitting the request body. @@ -133,22 +135,33 @@ class CallServerInterceptor( throw ProtocolException("Unexpected $HTTP_SWITCHING_PROTOCOLS code on HTTP/2 connection") } - val isUpgradeResponse = "upgrade".equals(response.header("Connection"), ignoreCase = true) - // TODO(jwilson): maybe call exchange.noRequestBody() if the upgrade failed? - response = - if (isUpgradeCode && isUpgradeRequest && isUpgradeResponse) { - 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.upgradeToSocket()) - .build() + val isUpgradeResponse = isUpgradeCode + && "upgrade".equals(response.header("Connection"), ignoreCase = true) + + response = when { + // This is an HTTP/1 upgrade. (This case includes web socket upgrades.) + isUpgradeRequest && isUpgradeResponse -> { + response + .newBuilder() + .body( + UnreadableResponseBody( + response.body.contentType(), + response.body.contentLength(), + ) + ) + .apply { + if (!forWebSocket) { + socket(exchange.upgradeToSocket()) + } + } + .build() + } + + // This is not an upgrade response. + else -> { + if (isUpgradeRequest) { + exchange.noRequestBody() // Failed upgrade request has no outbound data. } - } else { val responseBody = exchange.openResponseBody(response) response .newBuilder() @@ -167,6 +180,7 @@ class CallServerInterceptor( }, ).build() } + } if ("close".equals(response.request.header("Connection"), ignoreCase = true) || "close".equals(response.header("Connection"), ignoreCase = true) ) { diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index e3efab64e3cb..87cd5df4911d 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -17,6 +17,7 @@ package okhttp3.internal.http import assertk.assertThat import assertk.assertions.containsExactly +import assertk.assertions.hasMessage import assertk.assertions.isEqualTo import assertk.assertions.isNull import assertk.assertions.isTrue @@ -30,6 +31,7 @@ import okhttp3.Protocol import okhttp3.RecordingEventListener import okhttp3.RecordingHostnameVerifier import okhttp3.Request +import okhttp3.RequestBody.Companion.toRequestBody import okhttp3.internal.duplex.MockSocketHandler import okhttp3.testing.PlatformRule import okio.ProtocolException @@ -110,7 +112,6 @@ class HttpUpgradesTest { .Builder() .url(server.url("/")) .header("Connection", "upgrade") - .header("Upgrade", "tcp") .build() client.newCall(requestWithUpgrade).execute().use { response -> assertThat(response.code).isEqualTo(200) @@ -129,7 +130,6 @@ class HttpUpgradesTest { .Builder() .url(server.url("/")) .header("Connection", "upgrade") - .header("Upgrade", "tcp") .build() assertFailsWith { client.newCall(requestWithUpgrade).execute() @@ -176,6 +176,18 @@ class HttpUpgradesTest { ) } + @Test + fun upgradeRequestMustNotHaveABody() { + val e = assertFailsWith { + Request.Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .post("Hello".toRequestBody()) + .build() + } + assertThat(e).hasMessage("expected a null request body with 'Connection: upgrade'") + } + private fun enableTls(vararg protocols: Protocol) { client = client @@ -197,8 +209,6 @@ class HttpUpgradesTest { headersOf( "Connection", "upgrade", - "Upgrade", - "tcp", ), ) @@ -207,7 +217,6 @@ class HttpUpgradesTest { .Builder() .code(HTTP_SWITCHING_PROTOCOLS) .addHeader("Connection", "upgrade") - .addHeader("Upgrade", "tcp") .socketHandler(this) .build() } From 713683dec327d0a69234e20392d9513f49f04b19 Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Mon, 28 Jul 2025 10:21:24 -0400 Subject: [PATCH 16/17] Upgrades don't promise a response body --- .../kotlin/mockwebserver3/MockWebServer.kt | 15 ++++----- .../okhttp3/internal/http/HttpHeaders.kt | 2 +- .../okhttp3/internal/http/HttpUpgradesTest.kt | 33 +++++++++++++++++++ 3 files changed, 40 insertions(+), 10 deletions(-) diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 9d3921cab7ac..9117c2bc6802 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -601,18 +601,15 @@ public class MockWebServer : Closeable { } var reuseSocket = true - val requestWantsWebSockets = - "Upgrade".equals(request.headers["Connection"], ignoreCase = true) && + val requestWantsSocket = "Upgrade".equals(request.headers["Connection"], ignoreCase = true) + val requestWantsWebSocket = requestWantsSocket && "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) { + val responseWantsSocket = response.socketHandler != null + val responseWantsWebSocket = response.webSocketListener != null + if (requestWantsWebSocket && responseWantsWebSocket) { handleWebSocketUpgrade(socket, source, sink, request, response) reuseSocket = false - } else if (requestWantsTcp && responseWantsStream) { + } else if (requestWantsSocket && responseWantsSocket) { writeHttpResponse(socket, sink, response) reuseSocket = false } else { diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/HttpHeaders.kt index a47546399eb7..4df249a2dd49 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 == HTTP_SWITCHING_PROTOCOLS || responseCode >= 200) && + if ((responseCode < HTTP_CONTINUE || responseCode >= 200) && responseCode != HTTP_NO_CONTENT && responseCode != HTTP_NOT_MODIFIED ) { diff --git a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index 87cd5df4911d..c3b4ed4a3b64 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -118,6 +118,26 @@ class HttpUpgradesTest { assertThat(response.socket).isNull() assertThat(response.body.string()).isEqualTo("normal request") } + // Confirm there's no RequestBodyStart/RequestBodyEnd on failed upgrades. + assertThat(listener.recordedEventTypes()).containsExactly( + "CallStart", + "ProxySelectStart", + "ProxySelectEnd", + "DnsStart", + "DnsEnd", + "ConnectStart", + "ConnectEnd", + "ConnectionAcquired", + "RequestHeadersStart", + "RequestHeadersEnd", + "ResponseHeadersStart", + "ResponseHeadersEnd", + "FollowUpDecision", + "ResponseBodyStart", + "ResponseBodyEnd", + "ConnectionReleased", + "CallEnd", + ) } @Test @@ -149,6 +169,19 @@ class HttpUpgradesTest { assertThat(server.takeRequest().connectionIndex).isEqualTo(0) } + @Test + fun cannotReuseConnectionAfterUpgrade() { + upgrade() + + server.enqueue(MockResponse(body = "normal request")) + client.newCall(Request(server.url("/"))).execute().use { response -> + assertThat(response.body.string()).isEqualTo("normal request") + } + + assertThat(server.takeRequest().connectionIndex).isEqualTo(0) + assertThat(server.takeRequest().connectionIndex).isEqualTo(1) + } + @Test fun upgradeEvents() { upgrade() From d9d9122ce61e5823be46ad7ab9753edb57f1a64e Mon Sep 17 00:00:00 2001 From: Jesse Wilson Date: Mon, 28 Jul 2025 10:25:46 -0400 Subject: [PATCH 17/17] SpotlessApply --- .../kotlin/mockwebserver3/MockWebServer.kt | 3 +- .../internal/http/CallServerInterceptor.kt | 91 ++++++++++--------- .../okhttp3/internal/http/HttpUpgradesTest.kt | 16 ++-- 3 files changed, 57 insertions(+), 53 deletions(-) diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 9117c2bc6802..ece98492d3ee 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -602,7 +602,8 @@ public class MockWebServer : Closeable { var reuseSocket = true val requestWantsSocket = "Upgrade".equals(request.headers["Connection"], ignoreCase = true) - val requestWantsWebSocket = requestWantsSocket && + val requestWantsWebSocket = + requestWantsSocket && "websocket".equals(request.headers["Upgrade"], ignoreCase = true) val responseWantsSocket = response.socketHandler != null val responseWantsWebSocket = response.webSocketListener != null diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt index 7f017edbee30..4822d77d4bcb 100644 --- a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt +++ b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/http/CallServerInterceptor.kt @@ -43,8 +43,9 @@ class CallServerInterceptor( var responseBuilder: Response.Builder? = null var sendRequestException: IOException? = null val hasRequestBody = HttpMethod.permitsRequestBody(request.method) && requestBody != null - val isUpgradeRequest = !hasRequestBody - && "upgrade".equals(request.header("Connection"), ignoreCase = true) + val isUpgradeRequest = + !hasRequestBody && + "upgrade".equals(request.header("Connection"), ignoreCase = true) try { exchange.writeRequestHeaders(request) @@ -135,52 +136,52 @@ class CallServerInterceptor( throw ProtocolException("Unexpected $HTTP_SWITCHING_PROTOCOLS code on HTTP/2 connection") } - val isUpgradeResponse = isUpgradeCode - && "upgrade".equals(response.header("Connection"), ignoreCase = true) - - response = when { - // This is an HTTP/1 upgrade. (This case includes web socket upgrades.) - isUpgradeRequest && isUpgradeResponse -> { - response - .newBuilder() - .body( - UnreadableResponseBody( - response.body.contentType(), - response.body.contentLength(), - ) - ) - .apply { - if (!forWebSocket) { - socket(exchange.upgradeToSocket()) - } - } - .build() - } - - // This is not an upgrade response. - else -> { - if (isUpgradeRequest) { - exchange.noRequestBody() // Failed upgrade request has no outbound data. + val isUpgradeResponse = + isUpgradeCode && + "upgrade".equals(response.header("Connection"), ignoreCase = true) + + response = + when { + // This is an HTTP/1 upgrade. (This case includes web socket upgrades.) + isUpgradeRequest && isUpgradeResponse -> { + response + .newBuilder() + .body( + UnreadableResponseBody( + response.body.contentType(), + response.body.contentLength(), + ), + ).apply { + if (!forWebSocket) { + socket(exchange.upgradeToSocket()) + } + }.build() } - 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() + + // This is not an upgrade response. + else -> { + if (isUpgradeRequest) { + exchange.noRequestBody() // Failed upgrade request has no outbound data. + } + 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?!") } - 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/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt index c3b4ed4a3b64..a9666fc396b3 100644 --- a/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -211,13 +211,15 @@ class HttpUpgradesTest { @Test fun upgradeRequestMustNotHaveABody() { - val e = assertFailsWith { - Request.Builder() - .url(server.url("/")) - .header("Connection", "upgrade") - .post("Hello".toRequestBody()) - .build() - } + val e = + assertFailsWith { + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .post("Hello".toRequestBody()) + .build() + } assertThat(e).hasMessage("expected a null request body with 'Connection: upgrade'") }