diff --git a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt index 0f6f4381f282..ece98492d3ee 100644 --- a/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt +++ b/mockwebserver/src/main/kotlin/mockwebserver3/MockWebServer.kt @@ -89,6 +89,7 @@ import okio.BufferedSource import okio.ByteString import okio.Sink import okio.Timeout +import okio.asOkioSocket import okio.buffer import okio.sink import okio.source @@ -600,13 +601,18 @@ 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 - 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 (requestWantsSocket && responseWantsSocket) { + writeHttpResponse(socket, sink, response) + reuseSocket = false } else { writeHttpResponse(socket, sink, response) } @@ -865,6 +871,11 @@ public class MockWebServer : Closeable { writeHeaders(sink, response.headers) + if (response.socketHandler != null) { + response.socketHandler.handle(socket.asOkioSocket()) + return + } + val body = response.body ?: return socket.sleepWhileOpen(response.bodyDelayNanos) val responseBodySink = diff --git a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt index 956a0e824f59..7dd531c4ee2e 100644 --- a/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt +++ b/okhttp-testing-support/src/main/kotlin/okhttp3/internal/duplex/MockSocketHandler.kt @@ -82,7 +82,7 @@ class MockSocketHandler : SocketHandler { @JvmOverloads fun sendResponse( s: String, - responseSent: CountDownLatch = CountDownLatch(0), + responseSent: CountDownLatch = CountDownLatch(1), ) = apply { actions += { stream -> stream.sink.writeUtf8(s) diff --git a/okhttp/api/android/okhttp.api b/okhttp/api/android/okhttp.api index b1ee782b62ac..e0c2d0e41b89 100644 --- a/okhttp/api/android/okhttp.api +++ b/okhttp/api/android/okhttp.api @@ -1120,6 +1120,7 @@ public final class okhttp3/Response : java/io/Closeable { public final fun receivedResponseAtMillis ()J public final fun request ()Lokhttp3/Request; public final fun sentRequestAtMillis ()J + public final fun socket ()Lokio/Socket; public fun toString ()Ljava/lang/String; public final fun trailers ()Lokhttp3/Headers; } @@ -1142,6 +1143,7 @@ public class okhttp3/Response$Builder { public fun removeHeader (Ljava/lang/String;)Lokhttp3/Response$Builder; public fun request (Lokhttp3/Request;)Lokhttp3/Response$Builder; public fun sentRequestAtMillis (J)Lokhttp3/Response$Builder; + public fun socket (Lokio/Socket;)Lokhttp3/Response$Builder; public fun trailers (Lokhttp3/TrailersSource;)Lokhttp3/Response$Builder; } diff --git a/okhttp/api/jvm/okhttp.api b/okhttp/api/jvm/okhttp.api index b1ee782b62ac..e0c2d0e41b89 100644 --- a/okhttp/api/jvm/okhttp.api +++ b/okhttp/api/jvm/okhttp.api @@ -1120,6 +1120,7 @@ public final class okhttp3/Response : java/io/Closeable { public final fun receivedResponseAtMillis ()J public final fun request ()Lokhttp3/Request; public final fun sentRequestAtMillis ()J + public final fun socket ()Lokio/Socket; public fun toString ()Ljava/lang/String; public final fun trailers ()Lokhttp3/Headers; } @@ -1142,6 +1143,7 @@ public class okhttp3/Response$Builder { public fun removeHeader (Ljava/lang/String;)Lokhttp3/Response$Builder; public fun request (Lokhttp3/Request;)Lokhttp3/Response$Builder; public fun sentRequestAtMillis (J)Lokhttp3/Response$Builder; + public fun socket (Lokio/Socket;)Lokhttp3/Response$Builder; public fun trailers (Lokhttp3/TrailersSource;)Lokhttp3/Response$Builder; } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/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/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..89d30f22dc08 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,22 @@ class Exchange( ) } + fun upgradeToSocket(): Socket { + call.timeoutEarlyExit() + (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() { codec.carrier.noNewExchanges() } diff --git a/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt b/okhttp/src/commonJvmAndroid/kotlin/okhttp3/internal/connection/RealConnection.kt index a048f0cb81e9..f00dda51043c 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 @@ -294,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( @@ -312,6 +311,11 @@ class RealConnection internal constructor( } } + internal fun useAsSocket() { + socket.soTimeout = 0 + noNewExchanges() + } + 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..4822d77d4bcb 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,10 +42,14 @@ class CallServerInterceptor( var invokeStartEvent = true 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) 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. @@ -76,7 +80,7 @@ class CallServerInterceptor( exchange.noNewExchangesOnConnection() } } - } else { + } else if (!isUpgradeRequest) { exchange.noRequestBody() } @@ -127,28 +131,56 @@ 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") + } + + val isUpgradeResponse = + isUpgradeCode && + "upgrade".equals(response.header("Connection"), ignoreCase = true) + response = - if (forWebSocket && code == 101) { - // 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() - - override fun get(): Headers { - val source = responseBody.source() - if (source.isOpen) { - source.skipAll() - } - return peek() ?: error("null trailers after exhausting response body?!") + 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() + }.build() + } + + // 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?!") + } + }, + ).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/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/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/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 new file mode 100644 index 000000000000..a9666fc396b3 --- /dev/null +++ b/okhttp/src/jvmTest/kotlin/okhttp3/internal/http/HttpUpgradesTest.kt @@ -0,0 +1,257 @@ +/* + * 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.hasMessage +import assertk.assertions.isEqualTo +import assertk.assertions.isNull +import assertk.assertions.isTrue +import kotlin.test.assertFailsWith +import mockwebserver3.MockResponse +import mockwebserver3.MockWebServer +import mockwebserver3.junit5.StartStop +import okhttp3.Headers.Companion.headersOf +import okhttp3.OkHttpClientTestRule +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 +import okio.buffer +import okio.use +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.extension.RegisterExtension + +class HttpUpgradesTest { + @RegisterExtension + val platform = PlatformRule() + + @RegisterExtension + val clientTestRule = OkHttpClientTestRule() + + @StartStop + private val server = MockWebServer() + + private var listener = RecordingEventListener() + private val handshakeCertificates = platform.localhostHandshakeCertificates() + private var client = + clientTestRule + .newClientBuilder() + .eventListenerFactory(clientTestRule.wrap(listener)) + .build() + + @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() + } + } + + @Test + fun upgradeHttps() { + enableTls(Protocol.HTTP_1_1) + upgrade() + } + + @Test + fun upgradeRefusedByServer() { + server.enqueue(MockResponse(body = "normal request")) + val requestWithUpgrade = + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .build() + client.newCall(requestWithUpgrade).execute().use { response -> + assertThat(response.code).isEqualTo(200) + 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 + fun upgradeForbiddenOnHttp2() { + enableTls(Protocol.HTTP_2, Protocol.HTTP_1_1) + val socketHandler = MockSocketHandler() + server.enqueue(socketHandler.upgradeResponse()) + val requestWithUpgrade = + Request + .Builder() + .url(server.url("/")) + .header("Connection", "upgrade") + .build() + assertFailsWith { + client.newCall(requestWithUpgrade).execute() + } + } + + @Test + fun upgradesOnReusedConnection() { + server.enqueue(MockResponse(body = "normal request")) + client.newCall(Request(server.url("/"))).execute().use { response -> + assertThat(response.body.string()).isEqualTo("normal request") + } + + upgrade() + + assertThat(server.takeRequest().connectionIndex).isEqualTo(0) + 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() + + assertThat(listener.recordedEventTypes()).containsExactly( + "CallStart", + "ProxySelectStart", + "ProxySelectEnd", + "DnsStart", + "DnsEnd", + "ConnectStart", + "ConnectEnd", + "ConnectionAcquired", + "RequestHeadersStart", + "RequestHeadersEnd", + "ResponseHeadersStart", + "ResponseHeadersEnd", + "RequestBodyStart", + "FollowUpDecision", + "ResponseBodyStart", + "ResponseBodyEnd", + "RequestBodyEnd", + "ConnectionReleased", + "CallEnd", + ) + } + + @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 + .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", + ), + ) + + private fun MockSocketHandler.upgradeResponse() = + MockResponse + .Builder() + .code(HTTP_SWITCHING_PROTOCOLS) + .addHeader("Connection", "upgrade") + .socketHandler(this) + .build() +}