File
Blob: src/workerd/api/sockets.c++
| 1 | // Copyright (c) 2017-2022 Cloudflare, Inc. |
| 2 | // Licensed under the Apache 2.0 license found in the LICENSE file or at: |
| 3 | // https://opensource.org/licenses/Apache-2.0 |
| 4 | |
| 5 | #include "sockets.h" |
| 6 | |
| 7 | #include "global-scope.h" |
| 8 | #include "streams/standard.h" |
| 9 | #include "system-streams.h" |
| 10 | |
| 11 | #include <workerd/io/io-context.h> |
| 12 | #include <workerd/io/worker-interface.h> |
| 13 | #include <workerd/jsg/exception.h> |
| 14 | #include <workerd/jsg/url.h> |
| 15 | #include <workerd/util/autogate.h> |
| 16 | |
| 17 | namespace workerd::api { |
| 18 | |
| 19 | namespace { |
| 20 | |
| 21 | // This function performs some basic length and characters checks, it does not guarantee that |
| 22 | // the specified host is a valid domain. It should only be used to reject malicious |
| 23 | // hosts. |
| 24 | bool isValidHost(kj::StringPtr host) { |
| 25 | if (host.size() > 255 || host.size() == 0) { |
| 26 | // RFC1035 states that maximum domain name length is 255 octets. |
| 27 | // |
| 28 | // IP addresses are always shorter, so we take the max domain length instead. |
| 29 | return false; |
| 30 | } |
| 31 | |
| 32 | for (auto i: kj::indices(host)) { |
| 33 | switch (host[i]) { |
| 34 | case '-': |
| 35 | case '.': |
| 36 | case '_': |
| 37 | case '[': |
| 38 | case ']': |
| 39 | case ':': // For IPv6. |
| 40 | break; |
| 41 | default: |
| 42 | if ((host[i] >= 'a' && host[i] <= 'z') || (host[i] >= 'A' && host[i] <= 'Z') || |
| 43 | (host[i] >= '0' && host[i] <= '9')) { |
| 44 | break; |
| 45 | } |
| 46 | return false; |
| 47 | } |
| 48 | } |
| 49 | return true; |
| 50 | } |
| 51 | |
| 52 | SecureTransportKind parseSecureTransport(SocketOptions& opts) { |
| 53 | auto value = KJ_UNWRAP_OR_RETURN(opts.secureTransport, SecureTransportKind::OFF).begin(); |
| 54 | if (value == "off"_kj) { |
| 55 | return SecureTransportKind::OFF; |
| 56 | } else if (value == "starttls"_kj) { |
| 57 | return SecureTransportKind::STARTTLS; |
| 58 | } else if (value == "on"_kj) { |
| 59 | return SecureTransportKind::ON; |
| 60 | } else { |
| 61 | JSG_FAIL_REQUIRE( |
| 62 | TypeError, kj::str("Unsupported value in secureTransport socket option: ", value)); |
| 63 | } |
| 64 | } |
| 65 | |
| 66 | bool getAllowHalfOpen(jsg::Optional<SocketOptions>& opts) { |
| 67 | KJ_IF_SOME(o, opts) { |
| 68 | return o.allowHalfOpen; |
| 69 | } |
| 70 | |
| 71 | // The allowHalfOpen flag is false by default. |
| 72 | return false; |
| 73 | } |
| 74 | |
| 75 | kj::Maybe<uint64_t> getWritableHighWaterMark(jsg::Optional<SocketOptions>& opts) { |
| 76 | KJ_IF_SOME(o, opts) { |
| 77 | return o.highWaterMark; |
| 78 | } |
| 79 | return kj::none; |
| 80 | } |
| 81 | |
| 82 | } // namespace |
| 83 | |
| 84 | // Forward declarations |
| 85 | class StreamWorkerInterface; |
| 86 | |
| 87 | jsg::Ref<Socket> setupSocket(jsg::Lock& js, |
| 88 | kj::Own<kj::AsyncIoStream> connection, |
| 89 | kj::Maybe<kj::String> remoteAddress, |
| 90 | kj::Maybe<kj::String> localAddress, |
| 91 | jsg::Optional<SocketOptions> options, |
| 92 | kj::Own<kj::TlsStarterCallback> tlsStarter, |
| 93 | SecureTransportKind secureTransport, |
| 94 | kj::Maybe<kj::String> domain, |
| 95 | bool isDefaultFetchPort, |
| 96 | kj::Maybe<jsg::PromiseResolverPair<SocketInfo>> maybeOpenedPrPair) { |
| 97 | auto& ioContext = IoContext::current(); |
| 98 | |
| 99 | // Disconnection handling is annoyingly complicated: |
| 100 | // |
| 101 | // We can't just context.awaitIo(connection->whenWriteDisconnected()) directly, because the |
| 102 | // Socket could be GC'ed before `whenWriteDisconnected()` completes, causing the underlying |
| 103 | // `connection` to be destroyed. By KJ rules, we are required to cancel the promise returned by |
| 104 | // `whenWriteDisconnected()` before destroying `connection`. But there's no way to cancel a |
| 105 | // promise passed to `context.awaitIo()`. We have to hold the promise directly in `Socket`, so |
| 106 | // that we can cancel it on destruction. But we *do* want to create a JS promise that resolves |
| 107 | // on disconnect, which is what awaitIo() would give us. |
| 108 | // |
| 109 | // So, we have to chain through a promise/fulfiller pair. The `Socket` holds |
| 110 | // `watchForDisconnectTask`, which is a `kj::Promise<void>` representing a task that waits for |
| 111 | // `whenWriteDisconnected()` and then fulfills the fulfiller end of `disconnectedPaf` with |
| 112 | // `false`. If the task is canceled, we instead fulfill `disconnectedPaf` with `true`. |
| 113 | // |
| 114 | // We then use `context.awaitIo()` to await the promise end of `disconnectedPaf`, and this gives |
| 115 | // us our `closed` promise. Well, almost... |
| 116 | // |
| 117 | // There's another wrinkle: There are some circumstances where we want to resolve the `closed` |
| 118 | // promise directly from an API call. We'd rather this did not have to drop out of the isolate |
| 119 | // and enter it a gain. So, our `awaitIo()` actually awaits a task that listens for the |
| 120 | // disconnected promise and then resolves some other JS resolver, `closedResolver`. |
| 121 | auto disconnectedPaf = kj::newPromiseAndFulfiller<bool>(); |
| 122 | auto& disconnectedFulfiller = *disconnectedPaf.fulfiller; |
| 123 | auto deferredCancelDisconnected = |
| 124 | kj::defer([fulfiller = kj::mv(disconnectedPaf.fulfiller)]() mutable { |
| 125 | // In case the `whenWriteDisconected()` listener task is canceled without fulfilling the |
| 126 | // fulfiller, we want to silently fulfill it. This will happen when the Socket is GC'ed. |
| 127 | fulfiller->fulfill(true); |
| 128 | }); |
| 129 | |
| 130 | static auto constexpr handleDisconnected = |
| 131 | [](kj::AsyncIoStream& connection, |
| 132 | kj::PromiseFulfiller<bool>& fulfiller) -> kj::Promise<void> { |
| 133 | try { |
| 134 | co_await connection.whenWriteDisconnected(); |
| 135 | fulfiller.fulfill(false); |
| 136 | } catch (...) { |
| 137 | auto exception = kj::getCaughtExceptionAsKj(); |
| 138 | fulfiller.reject(kj::mv(exception)); |
| 139 | } |
| 140 | }; |
| 141 | |
| 142 | auto watchForDisconnectTask = handleDisconnected(*connection, disconnectedFulfiller) |
| 143 | .attach(kj::mv(deferredCancelDisconnected)); |
| 144 | |
| 145 | auto closedPrPair = js.newPromiseAndResolver<void>(); |
| 146 | closedPrPair.promise.markAsHandled(js); |
| 147 | |
| 148 | ioContext.awaitIo(js, kj::mv(disconnectedPaf.promise)) |
| 149 | .then( |
| 150 | js, [resolver = closedPrPair.resolver.addRef(js)](jsg::Lock& js, bool canceled) mutable { |
| 151 | // We want to silently ignore the canceled case, without ever resolving anything. Note that |
| 152 | // if the application actually fetches the `closed` promise, then the JSG glue will prevent |
| 153 | // the socket from being GC'ed until that promise resolves, so it won't be canceled. |
| 154 | if (!canceled) { |
| 155 | resolver.resolve(js); |
| 156 | } |
| 157 | }, [resolver = closedPrPair.resolver.addRef(js)](jsg::Lock& js, jsg::Value exception) mutable { |
| 158 | resolver.reject(js, exception.getHandle(js)); |
| 159 | }); |
| 160 | |
| 161 | auto refcountedConnection = kj::refcountedWrapper(kj::mv(connection)); |
| 162 | // Initialize the readable/writable streams with the readable/writable sides of an AsyncIoStream. |
| 163 | auto sysStreams = newSystemMultiStream(*refcountedConnection, ioContext); |
| 164 | auto readable = js.alloc<ReadableStream>(ioContext, kj::mv(sysStreams.readable)); |
| 165 | auto allowHalfOpen = getAllowHalfOpen(options); |
| 166 | kj::Maybe<jsg::Promise<void>> eofPromise; |
| 167 | if (!allowHalfOpen) { |
| 168 | eofPromise = readable->onEof(js); |
| 169 | } |
| 170 | auto openedPrPair = kj::mv(maybeOpenedPrPair).orDefault(js.newPromiseAndResolver<SocketInfo>()); |
| 171 | openedPrPair.promise.markAsHandled(js); |
| 172 | auto writable = js.alloc<WritableStream>(ioContext, kj::mv(sysStreams.writable), |
| 173 | ioContext.getMetrics().tryCreateWritableByteStreamObserver(), |
| 174 | getWritableHighWaterMark(options), openedPrPair.promise.whenResolved(js)); |
| 175 | |
| 176 | auto result = js.alloc<Socket>(js, ioContext, kj::mv(refcountedConnection), kj::mv(remoteAddress), |
| 177 | kj::mv(localAddress), kj::mv(readable), kj::mv(writable), kj::mv(closedPrPair), |
| 178 | kj::mv(watchForDisconnectTask), kj::mv(options), kj::mv(tlsStarter), secureTransport, |
| 179 | kj::mv(domain), isDefaultFetchPort, kj::mv(openedPrPair)); |
| 180 | |
| 181 | KJ_IF_SOME(p, eofPromise) { |
| 182 | result->handleReadableEof(js, kj::mv(p)); |
| 183 | } |
| 184 | return result; |
| 185 | } |
| 186 | |
| 187 | jsg::Ref<Socket> connectImplNoOutputLock(jsg::Lock& js, |
| 188 | kj::Maybe<jsg::Ref<Fetcher>> fetcher, |
| 189 | AnySocketAddress address, |
| 190 | jsg::Optional<SocketOptions> options) { |
| 191 | |
| 192 | auto& ioContext = IoContext::current(); |
| 193 | |
| 194 | // Extract the domain/ip we are connecting to from the address. |
| 195 | kj::String domain; |
| 196 | bool isDefaultFetchPort = false; |
| 197 | |
| 198 | KJ_SWITCH_ONEOF(address) { |
| 199 | KJ_CASE_ONEOF(str, kj::String) { |
| 200 | // We need just the hostname part of the address, i.e. we want to strip out the port. |
| 201 | // We do this using the standard URL parser since it will handle IPv6 for us as well. |
| 202 | auto input = kj::str("fake://", str); |
| 203 | auto url = JSG_REQUIRE_NONNULL( |
| 204 | jsg::Url::tryParse(input.asPtr()), TypeError, "Specified address could not be parsed."); |
| 205 | auto host = url.getHostname(); |
| 206 | auto port = url.getPort(); |
| 207 | JSG_REQUIRE(host != ""_kj, TypeError, "Specified address is missing hostname."); |
| 208 | JSG_REQUIRE(port != ""_kj, TypeError, "Specified address is missing port."); |
| 209 | isDefaultFetchPort = port == "443"_kj || port == "80"_kj; |
| 210 | domain = kj::str(host); |
| 211 | } |
| 212 | KJ_CASE_ONEOF(record, SocketAddress) { |
| 213 | domain = kj::heapString(record.hostname); |
| 214 | isDefaultFetchPort = record.port == 443 || record.port == 80; |
| 215 | } |
| 216 | } |
| 217 | |
| 218 | // Convert the address to a string that we can pass to kj. |
| 219 | auto addressStr = kj::str(""); |
| 220 | KJ_SWITCH_ONEOF(address) { |
| 221 | KJ_CASE_ONEOF(str, kj::String) { |
| 222 | addressStr = kj::mv(str); |
| 223 | } |
| 224 | KJ_CASE_ONEOF(record, SocketAddress) { |
| 225 | addressStr = kj::str(record.hostname, ":", record.port); |
| 226 | } |
| 227 | } |
| 228 | |
| 229 | JSG_REQUIRE(isValidHost(addressStr), TypeError, |
| 230 | "Specified address is empty string, contains unsupported characters or is too long."); |
| 231 | |
| 232 | jsg::Ref<Fetcher> actualFetcher = nullptr; |
| 233 | KJ_IF_SOME(f, fetcher) { |
| 234 | actualFetcher = kj::mv(f); |
| 235 | } else { |
| 236 | // Support calling into arbitrary callbacks for any registered "magic" addresses for which |
| 237 | // custom connect() logic is needed. Note that these overrides should only apply to calls of the |
| 238 | // global connect() method, not for fetcher->connect(), hence why we check for them here. |
| 239 | KJ_IF_SOME(fn, ioContext.getCurrentLock().getGlobalScope().getConnectOverride(addressStr)) { |
| 240 | return fn(js); |
| 241 | } |
| 242 | actualFetcher = |
| 243 | js.alloc<Fetcher>(IoContext::NULL_CLIENT_CHANNEL, Fetcher::RequiresHostAndProtocol::YES); |
| 244 | } |
| 245 | |
| 246 | CfProperty cf; |
| 247 | kj::Own<WorkerInterface> client = |
| 248 | actualFetcher->getClient(ioContext, cf.serialize(js), "connect"_kjc); |
| 249 | |
| 250 | // Set up the connection. |
| 251 | auto headers = kj::heap<kj::HttpHeaders>(ioContext.getHeaderTable()); |
| 252 | kj::HttpConnectSettings httpConnectSettings = {.useTls = false}; |
| 253 | SecureTransportKind secureTransport = SecureTransportKind::OFF; |
| 254 | KJ_IF_SOME(opts, options) { |
| 255 | secureTransport = parseSecureTransport(opts); |
| 256 | httpConnectSettings.useTls = secureTransport == SecureTransportKind::ON; |
| 257 | } |
| 258 | kj::Own<kj::TlsStarterCallback> tlsStarter = kj::heap<kj::TlsStarterCallback>(); |
| 259 | httpConnectSettings.tlsStarter = tlsStarter; |
| 260 | |
| 261 | KJ_IF_SOME(promise, |
| 262 | util::Autogate::isEnabled(util::AutogateKey::TCP_SOCKET_CONNECT_OUTPUT_GATE) |
| 263 | ? ioContext.waitForOutputLocksIfNecessary() |
| 264 | : kj::none) { |
| 265 | // Wrap the real WorkerInterface in a promised interface that defers connect |
| 266 | // until the DO output gate clears. |
| 267 | client = newPromisedWorkerInterface( |
| 268 | kj::mv(promise).then([client = kj::mv(client)]() mutable { return kj::mv(client); })); |
| 269 | } |
| 270 | |
| 271 | auto httpClient = asHttpClient(kj::mv(client)); |
| 272 | auto request = httpClient->connect(addressStr, *headers, httpConnectSettings); |
| 273 | request.connection = request.connection.attach(kj::mv(httpClient)); |
| 274 | |
| 275 | auto result = setupSocket(js, kj::mv(request.connection), kj::mv(addressStr), |
| 276 | kj::none /* localAddress */, kj::mv(options), kj::mv(tlsStarter), secureTransport, |
| 277 | kj::mv(domain), isDefaultFetchPort, kj::none /* maybeOpenedPrPair */); |
| 278 | // `handleProxyStatus` needs an initialized refcount to use `JSG_THIS`, hence it cannot be |
| 279 | // called in Socket's constructor. Also it's only necessary when creating a Socket as a result of |
| 280 | // a `connect`. |
| 281 | result->handleProxyStatus(js, kj::mv(request.status)); |
| 282 | return result; |
| 283 | } |
| 284 | |
| 285 | jsg::Ref<Socket> connectImpl(jsg::Lock& js, |
| 286 | kj::Maybe<jsg::Ref<Fetcher>> fetcher, |
| 287 | AnySocketAddress address, |
| 288 | jsg::Optional<SocketOptions> options) { |
| 289 | // When the TCP_SOCKET_CONNECT_OUTPUT_GATE autogate is enabled, the output gate wait is |
| 290 | // handled inside connectImplNoOutputLock via a deferred connect task, so no separate wait |
| 291 | // is needed here. TODO(cleanup): rename connectImplNoOutputLock once the autogate is removed. |
| 292 | return connectImplNoOutputLock(js, kj::mv(fetcher), kj::mv(address), kj::mv(options)); |
| 293 | } |
| 294 | |
| 295 | jsg::Promise<void> Socket::close(jsg::Lock& js) { |
| 296 | if (isClosing) { |
| 297 | return closedPromiseCopy.whenResolved(js); |
| 298 | } |
| 299 | |
| 300 | isClosing = true; |
| 301 | writable->getController().setPendingClosure(); |
| 302 | readable->getController().setPendingClosure(); |
| 303 | |
| 304 | // Wait until the socket connects (successfully or otherwise) |
| 305 | return openedPromiseCopy.whenResolved(js) |
| 306 | .then(js, |
| 307 | [this](jsg::Lock& js) { |
| 308 | if (!writable->getController().isClosedOrClosing()) { |
| 309 | return writable->getController().flush(js); |
| 310 | } else { |
| 311 | return js.resolvedPromise(); |
| 312 | } |
| 313 | }) |
| 314 | .then(js, |
| 315 | [this](jsg::Lock& js) { |
| 316 | // Forcibly abort the readable/writable streams. |
| 317 | auto cancelPromise = readable->getController().cancel(js, kj::none); |
| 318 | auto abortPromise = writable->getController().abort(js, kj::none); |
| 319 | |
| 320 | // The below is effectively `Promise.all(cancelPromise, abortPromise)` |
| 321 | return cancelPromise.then(js, [abortPromise = kj::mv(abortPromise)](jsg::Lock& js) mutable { |
| 322 | return kj::mv(abortPromise); |
| 323 | }); |
| 324 | }) |
| 325 | .then(js, [this](jsg::Lock& js) { |
| 326 | // Destroy the connection stream to close the connection. |
| 327 | { auto _ = kj::mv(connectionData); } |
| 328 | connectionData = kj::none; |
| 329 | |
| 330 | resolveFulfiller(js, kj::none); |
| 331 | return js.resolvedPromise(); |
| 332 | }).catch_(js, [this](jsg::Lock& js, jsg::Value err) { errorHandler(js, kj::mv(err)); }); |
| 333 | } |
| 334 | |
| 335 | jsg::Ref<Socket> Socket::startTls(jsg::Lock& js, jsg::Optional<TlsOptions> tlsOptions) { |
| 336 | JSG_REQUIRE( |
| 337 | secureTransport != SecureTransportKind::ON, TypeError, "Cannot startTls on a TLS socket."); |
| 338 | JSG_REQUIRE(connectionData != kj::none, TypeError, |
| 339 | "The connection was closed before startTls could be started."); |
| 340 | auto invalidOptKindMsg = |
| 341 | "The `secureTransport` socket option must be set to 'starttls' for startTls to be used."; |
| 342 | JSG_REQUIRE(secureTransport == SecureTransportKind::STARTTLS, TypeError, invalidOptKindMsg); |
| 343 | JSG_REQUIRE(domain != kj::none, TypeError, "startTls can only be called once."); |
| 344 | |
| 345 | // The current socket's writable buffers need to be flushed. The socket's WritableStream is backed |
| 346 | // by an AsyncIoStream which doesn't implement any buffering, so we don't need to worry about |
| 347 | // flushing. But the JS WritableStream holds a queue so some data may still be buffered. This |
| 348 | // means we need to flush the WritableStream. |
| 349 | // |
| 350 | // Detach the AsyncIoStream from the Writable/Readable streams and make them unusable. |
| 351 | auto& context = IoContext::current(); |
| 352 | auto openedPrPair = js.newPromiseAndResolver<SocketInfo>(); |
| 353 | auto secureStreamPromise = context.awaitJs(js, |
| 354 | writable->flush(js).then(js, |
| 355 | // The openedResolver is a jsg::Promise::Resolver. It should be gc visited here in |
| 356 | // case the opened promise it resolves captures a circular references to itself in |
| 357 | // JavaScript (which is most likely). This prevents a possible memory leak. |
| 358 | // We also capture a strong reference to the original Socket instance that is being |
| 359 | // upgraded in order to prevent it from being GC'd while we are waiting for the |
| 360 | // flush to complete. While it is unlikely to be GC'd while we are waiting because |
| 361 | // the user code *likely* is holding a active reference to it at this point, we |
| 362 | // don't want to take any chances. This prevents a possible UAF. |
| 363 | JSG_VISITABLE_LAMBDA((self = JSG_THIS, domain = kj::heapString(KJ_ASSERT_NONNULL(domain)), |
| 364 | tlsOptions = kj::mv(tlsOptions), |
| 365 | openedResolver = openedPrPair.resolver.addRef(js), |
| 366 | remoteAddress = mapCopyString(remoteAddress), |
| 367 | localAddress = mapCopyString(localAddress)), |
| 368 | (self, openedResolver), (jsg::Lock & js) mutable { |
| 369 | auto& context = IoContext::current(); |
| 370 | |
| 371 | self->writable->detach(js); |
| 372 | self->readable->detach(js, true); |
| 373 | |
| 374 | // We should set this before closedResolver.resolve() in order to give the user |
| 375 | // the option to check if the closed promise is resolved due to upgrade or not. |
| 376 | self->upgraded = true; |
| 377 | self->closedResolver.resolve(js); |
| 378 | |
| 379 | auto acceptedHostname = domain.asPtr(); |
| 380 | KJ_IF_SOME(s, tlsOptions) { |
| 381 | KJ_IF_SOME(expectedHost, s.expectedServerHostname) { |
| 382 | acceptedHostname = expectedHost; |
| 383 | } else { |
| 384 | } // Needed to avoid compiler error/warning |
| 385 | } else { |
| 386 | } // Needed to avoid compiler error/warning |
| 387 | |
| 388 | // All non-secure sockets should have `connectionData` with a `tlsStarter`. |
| 389 | // Though since it's inside an IoOwn, if the request's IoContext has ended |
| 390 | // then `connectionData` will be null. This can happen if the flush operation is taking |
| 391 | // a particularly long time (EW-8538), so we throw a JSG error if that's the case. |
| 392 | auto& connData = JSG_REQUIRE_NONNULL(self->connectionData, TypeError, |
| 393 | "The connection was closed before startTls completed."); |
| 394 | |
| 395 | auto& tlsStarter = connData->tlsStarter; |
| 396 | |
| 397 | // Fork the starter promise because we need to create two separate things waiting |
| 398 | // on it below. The first is resolving the openedResolver with a JS promise that |
| 399 | // wraps one branch, the second is the kj::Promise that we use to resolve the |
| 400 | // secureStream for the promised stream. This keeps us from having to bounce in and |
| 401 | // out of the JS isolate lock. |
| 402 | auto forkedPromise = KJ_ASSERT_NONNULL(*tlsStarter)(acceptedHostname).fork(); |
| 403 | |
| 404 | openedResolver.resolve(js, |
| 405 | context.awaitIo(js, forkedPromise.addBranch(), |
| 406 | [remoteAddress = kj::mv(remoteAddress), |
| 407 | localAddress = kj::mv(localAddress)]( |
| 408 | jsg::Lock& js) mutable -> SocketInfo { |
| 409 | return SocketInfo{ |
| 410 | .remoteAddress = kj::mv(remoteAddress), |
| 411 | .localAddress = kj::mv(localAddress), |
| 412 | }; |
| 413 | })); |
| 414 | |
| 415 | // Move the stream out of the plain text socket, to ensure the stream is properly |
| 416 | // destroyed when the socket is closed. |
| 417 | kj::Own<kj::AsyncIoStream> stream = connData->connectionStream->addWrappedRef(); |
| 418 | self->connectionData = kj::none; |
| 419 | |
| 420 | auto secureStream = forkedPromise.addBranch().then( |
| 421 | [stream = kj::mv(stream)]() mutable { return kj::mv(stream); }); |
| 422 | |
| 423 | return kj::newPromisedStream(kj::mv(secureStream)); |
| 424 | }))); |
| 425 | |
| 426 | // The existing tlsStarter gets consumed and we won't need it again. Pass in an empty tlsStarter |
| 427 | // to `setupSocket`. |
| 428 | auto newTlsStarter = kj::heap<kj::TlsStarterCallback>(); |
| 429 | return setupSocket(js, kj::newPromisedStream(kj::mv(secureStreamPromise)), |
| 430 | mapCopyString(remoteAddress), mapCopyString(localAddress), kj::mv(options), |
| 431 | kj::mv(newTlsStarter), SecureTransportKind::ON, kj::mv(domain), isDefaultFetchPort, |
| 432 | kj::mv(openedPrPair)); |
| 433 | } |
| 434 | |
| 435 | void Socket::handleProxyStatus( |
| 436 | jsg::Lock& js, kj::Promise<kj::HttpClient::ConnectRequest::Status> status) { |
| 437 | auto& context = IoContext::current(); |
| 438 | auto errorHandler = [](kj::Exception&& e) { |
| 439 | // Let's not log errors when we have a disconnected exception. |
| 440 | // If we don't filter this out, whenever connect() fails, we'll |
| 441 | // have noisy errors even though the user catches the error on JS side. |
| 442 | if (e.getType() != kj::Exception::Type::DISCONNECTED) { |
| 443 | LOG_ERROR_PERIODICALLY("Socket proxy disconnected abruptly", e); |
| 444 | } |
| 445 | return kj::HttpClient::ConnectRequest::Status(500, nullptr, kj::Own<kj::HttpHeaders>()); |
| 446 | }; |
| 447 | auto func = [this, self = JSG_THIS]( |
| 448 | jsg::Lock& js, kj::HttpClient::ConnectRequest::Status&& status) -> void { |
| 449 | if (status.statusCode < 200 || status.statusCode >= 300) { |
| 450 | // If the status indicates an unsuccessful connection we need to reject the `closeFulfiller` |
| 451 | // with an exception. This will reject the socket's `closed` promise. |
| 452 | auto msg = kj::str("proxy request failed, cannot connect to the specified address"); |
| 453 | if (isDefaultFetchPort) { |
| 454 | msg = kj::str(msg, ". It looks like you might be trying to connect to a HTTP-based service", |
| 455 | " — consider using fetch instead"); |
| 456 | } else if (remoteAddress.orDefault(kj::String()).contains(".hyperdrive.local"_kj)) { |
| 457 | // No attempts to connect to Hyperdrive should end up here, since they go through the other |
| 458 | // version of handleProxyStatus. If they end up here somehow, log about it to get some |
| 459 | // context that can aid in debugging. |
| 460 | LOG_WARNING_PERIODICALLY( |
| 461 | "attempt to connect to Hyperdrive failed to trigger connectOverride", remoteAddress, |
| 462 | status.statusCode, status.statusText); |
| 463 | } |
| 464 | handleProxyError(js, JSG_KJ_EXCEPTION(FAILED, Error, msg)); |
| 465 | } else { |
| 466 | // For outbound sockets we have no useful local address to expose. Inbound sockets (produced |
| 467 | // by the `connect()` handler dispatch path) populate `localAddress` with the CONNECT |
| 468 | // authority that the peer targeted. |
| 469 | openedResolver.resolve(js, |
| 470 | SocketInfo{ |
| 471 | .remoteAddress = mapCopyString(remoteAddress), |
| 472 | .localAddress = mapCopyString(localAddress), |
| 473 | }); |
| 474 | } |
| 475 | }; |
| 476 | auto result = context.awaitIo(js, status.catch_(kj::mv(errorHandler)), kj::mv(func)); |
| 477 | result.markAsHandled(js); |
| 478 | } |
| 479 | |
| 480 | void Socket::handleProxyStatus(jsg::Lock& js, kj::Promise<kj::Maybe<kj::Exception>> connectResult) { |
| 481 | // It's kind of weird to take a promise that resolves to a Maybe<Exception> but we can't just use |
| 482 | // a Promise<void> and put our logic in the error handler because awaitIo doesn't provide the |
| 483 | // jsg::Lock for void promises or to errorFunc implementations, only non-void success callbacks, |
| 484 | // but we need the lock in our callback here. |
| 485 | // TODO(cleanup): Extend awaitIo to provide the jsg::Lock in more cases. |
| 486 | auto& context = IoContext::current(); |
| 487 | auto errorHandler = [](kj::Exception&& e) -> kj::Maybe<kj::Exception> { |
| 488 | LOG_ERROR_PERIODICALLY("Socket proxy disconnected abruptly", e); |
| 489 | return KJ_EXCEPTION(FAILED, "connectResult raised an error"); |
| 490 | }; |
| 491 | auto func = [this, self = JSG_THIS](jsg::Lock& js, kj::Maybe<kj::Exception> result) -> void { |
| 492 | if (result != kj::none) { |
| 493 | handleProxyError(js, JSG_KJ_EXCEPTION(FAILED, Error, "connection attempt failed")); |
| 494 | } else { |
| 495 | // For outbound sockets we have no useful local address to expose. Inbound sockets (produced |
| 496 | // by the `connect()` handler dispatch path) populate `localAddress` with the CONNECT |
| 497 | // authority that the peer targeted. |
| 498 | openedResolver.resolve(js, |
| 499 | SocketInfo{ |
| 500 | .remoteAddress = mapCopyString(remoteAddress), |
| 501 | .localAddress = mapCopyString(localAddress), |
| 502 | }); |
| 503 | } |
| 504 | }; |
| 505 | auto result = context.awaitIo(js, connectResult.catch_(kj::mv(errorHandler)), kj::mv(func)); |
| 506 | result.markAsHandled(js); |
| 507 | } |
| 508 | |
| 509 | void Socket::handleProxyError(jsg::Lock& js, kj::Exception e) { |
| 510 | resolveFulfiller(js, e.clone()); |
| 511 | openedResolver.reject(js, e.clone()); |
| 512 | readable->getController().cancel(js, kj::none).markAsHandled(js); |
| 513 | writable->getController().abort(js, js.error(e.getDescription())).markAsHandled(js); |
| 514 | } |
| 515 | |
| 516 | void Socket::handleReadableEof(jsg::Lock& js, jsg::Promise<void> onEof) { |
| 517 | KJ_ASSERT(!getAllowHalfOpen(options)); |
| 518 | // Listen for EOF on the ReadableStream. |
| 519 | onEof |
| 520 | .then( |
| 521 | js, |
| 522 | JSG_VISITABLE_LAMBDA( |
| 523 | (ref = JSG_THIS), (ref), (jsg::Lock& js) { return ref->maybeCloseWriteSide(js); })) |
| 524 | .markAsHandled(js); |
| 525 | } |
| 526 | |
| 527 | jsg::Promise<void> Socket::maybeCloseWriteSide(jsg::Lock& js) { |
| 528 | // When `allowHalfOpen` is set to true then we do not automatically close the write side on EOF. |
| 529 | // This code shouldn't even run since we don't set up a callback which calls it unless |
| 530 | // `allowHalfOpen` is false. |
| 531 | KJ_ASSERT(!getAllowHalfOpen(options)); |
| 532 | |
| 533 | // Do not call `close` on a controller that has already been closed or is in the process |
| 534 | // of closing. |
| 535 | if (writable->getController().isClosedOrClosing()) { |
| 536 | return js.resolvedPromise(); |
| 537 | } |
| 538 | |
| 539 | // We want to close the socket, but only after its WritableStream has been flushed. We do this |
| 540 | // below by calling `close` on the WritableStream which ensures that any data pending on it |
| 541 | // is flushed. Then once the `close` either completes or fails we can be sure that any data has |
| 542 | // been flushed. |
| 543 | return writable->getController() |
| 544 | .close(js) |
| 545 | .catch_(js, |
| 546 | JSG_VISITABLE_LAMBDA((ref = JSG_THIS), (ref), |
| 547 | (jsg::Lock& js, jsg::Value&& exc) { |
| 548 | ref->closedResolver.reject(js, exc.getHandle(js)); |
| 549 | })) |
| 550 | .then(js, JSG_VISITABLE_LAMBDA((ref = JSG_THIS), (ref), (jsg::Lock& js) { |
| 551 | ref->closedResolver.resolve(js); |
| 552 | })); |
| 553 | } |
| 554 | |
| 555 | jsg::Ref<Socket> SocketsModule::connect( |
| 556 | jsg::Lock& js, AnySocketAddress address, jsg::Optional<SocketOptions> options) { |
| 557 | return connectImpl(js, kj::none, kj::mv(address), kj::mv(options)); |
| 558 | } |
| 559 | |
| 560 | kj::Own<kj::AsyncIoStream> Socket::takeConnectionStream(jsg::Lock& js) { |
| 561 | // Set this so that if `close` is called after this, that no closure steps are taken and instead |
| 562 | // the `close` is a no-op. |
| 563 | isClosing = true; |
| 564 | |
| 565 | // We do not care if the socket was disturbed, we require the user to ensure the socket is not |
| 566 | // being used. |
| 567 | writable->detach(js); |
| 568 | readable->detach(js, true); |
| 569 | |
| 570 | // Move the stream out of the socket, to ensure the stream is properly destroyed when the |
| 571 | // caller is done with it. |
| 572 | auto& dataConn = JSG_REQUIRE_NONNULL( |
| 573 | connectionData, TypeError, "The socket connection is closed or was already taken."); |
| 574 | // Attach tlsStarter to the wrapper so it survives as long as the connection stream |
| 575 | // and is destroyed before the stream itself. |
| 576 | auto wrapper = dataConn->connectionStream->addWrappedRef().attach(kj::mv(dataConn->tlsStarter)); |
| 577 | connectionData = kj::none; |
| 578 | closedResolver.resolve(js); |
| 579 | return wrapper; |
| 580 | } |
| 581 | |
| 582 | // Implementation of the custom factory for creating WorkerInterface instances from a socket |
| 583 | class StreamOutgoingFactory final: public Fetcher::OutgoingFactory, public kj::Refcounted { |
| 584 | public: |
| 585 | StreamOutgoingFactory(kj::Own<kj::AsyncIoStream> stream, |
| 586 | kj::EntropySource& entropySource, |
| 587 | const kj::HttpHeaderTable& headerTable) |
| 588 | : stream(kj::mv(stream)), |
| 589 | httpClient( |
| 590 | kj::newHttpClient(headerTable, *this->stream, {.entropySource = entropySource})) {} |
| 591 | |
| 592 | kj::Own<WorkerInterface> newSingleUseClient(kj::Maybe<kj::String> cfStr) override; |
| 593 | |
| 594 | private: |
| 595 | kj::Own<kj::AsyncIoStream> stream; |
| 596 | kj::Own<kj::HttpClient> httpClient; |
| 597 | friend class StreamWorkerInterface; |
| 598 | }; |
| 599 | |
| 600 | // Definition of the StreamWorkerInterface class |
| 601 | class StreamWorkerInterface final: public WorkerInterface { |
| 602 | public: |
| 603 | StreamWorkerInterface(kj::Own<StreamOutgoingFactory> factory): factory(kj::mv(factory)) {} |
| 604 | |
| 605 | kj::Promise<void> request(kj::HttpMethod method, |
| 606 | kj::StringPtr url, |
| 607 | const kj::HttpHeaders& headers, |
| 608 | kj::AsyncInputStream& requestBody, |
| 609 | kj::HttpService::Response& response) override { |
| 610 | // Parse the URL to extract the path |
| 611 | auto parsedUrl = KJ_REQUIRE_NONNULL(kj::Url::tryParse(url, kj::Url::Context::HTTP_PROXY_REQUEST, |
| 612 | {.percentDecode = false, .allowEmpty = true}), |
| 613 | "invalid url", url); |
| 614 | |
| 615 | // We need to convert the URL from proxy format (full URL in request line) to host format |
| 616 | // (path in request line, hostname in Host header). |
| 617 | auto newHeaders = headers.cloneShallow(); |
| 618 | newHeaders.setPtr(kj::HttpHeaderId::HOST, parsedUrl.host); |
| 619 | auto noHostUrl = parsedUrl.toString(kj::Url::Context::HTTP_REQUEST); |
| 620 | |
| 621 | // Create a new HTTP service from the client |
| 622 | auto service = kj::newHttpService(*factory->httpClient); |
| 623 | |
| 624 | // Forward the request to the service |
| 625 | co_await service->request(method, noHostUrl, newHeaders, requestBody, response); |
| 626 | } |
| 627 | |
| 628 | kj::Promise<void> connect(kj::StringPtr host, |
| 629 | const kj::HttpHeaders& headers, |
| 630 | kj::AsyncIoStream& connection, |
| 631 | ConnectResponse& response, |
| 632 | kj::HttpConnectSettings settings) override { |
| 633 | JSG_FAIL_REQUIRE(TypeError, |
| 634 | "connect is not something that can be done on a fetcher converted from a socket"); |
| 635 | } |
| 636 | |
| 637 | kj::Promise<void> prewarm(kj::StringPtr url) override { |
| 638 | KJ_UNIMPLEMENTED("prewarm() not supported on StreamWorkerInterface"); |
| 639 | } |
| 640 | |
| 641 | kj::Promise<ScheduledResult> runScheduled(kj::Date scheduledTime, kj::StringPtr cron) override { |
| 642 | KJ_UNIMPLEMENTED("runScheduled() not supported on StreamWorkerInterface"); |
| 643 | } |
| 644 | |
| 645 | kj::Promise<AlarmResult> runAlarm(kj::Date scheduledTime, uint32_t retryCount) override { |
| 646 | KJ_UNIMPLEMENTED("runAlarm() not supported on StreamWorkerInterface"); |
| 647 | } |
| 648 | |
| 649 | kj::Promise<CustomEvent::Result> customEvent(kj::Own<CustomEvent> event) override { |
| 650 | return event->notSupported(); |
| 651 | } |
| 652 | |
| 653 | private: |
| 654 | kj::Own<StreamOutgoingFactory> factory; |
| 655 | }; |
| 656 | |
| 657 | kj::Own<WorkerInterface> StreamOutgoingFactory::newSingleUseClient(kj::Maybe<kj::String> cfStr) { |
| 658 | JSG_ASSERT(stream.get() != nullptr, Error, |
| 659 | "Fetcher created from internalNewHttpClient can only be used once"); |
| 660 | // Create a WorkerInterface that wraps the stream, routing through getSubrequestNoChecks to apply |
| 661 | // external memory adjustment for GC pressure. |
| 662 | return IoContext::current().getSubrequestNoChecks([&](auto& tracing, auto& channelFactory) { |
| 663 | return kj::heap<StreamWorkerInterface>(kj::addRef(*this)); |
| 664 | }, {.inHouse = false, .wrapMetrics = false}); |
| 665 | } |
| 666 | |
| 667 | jsg::Promise<jsg::Ref<Fetcher>> SocketsModule::internalNewHttpClient( |
| 668 | jsg::Lock& js, jsg::Ref<Socket> socket) { |
| 669 | |
| 670 | // TODO(soon) check for nothing to read, this will require things using a promise so this function |
| 671 | // must remain returning a jsg::Promise waiting on a TODO for releaseLock |
| 672 | |
| 673 | // Flush the writable stream before taking the connection stream to ensure all data is written |
| 674 | // before the stream is detatched |
| 675 | return socket->getWritable()->flush(js).then( |
| 676 | js, JSG_VISITABLE_LAMBDA((socket = kj::mv(socket)), (socket), (jsg::Lock & js) mutable { |
| 677 | auto& ioctx = IoContext::current(); |
| 678 | |
| 679 | // Create our custom factory that will create client instances from this socket |
| 680 | kj::Own<Fetcher::OutgoingFactory> outgoingFactory = kj::refcounted<StreamOutgoingFactory>( |
| 681 | socket->takeConnectionStream(js), ioctx.getEntropySource(), ioctx.getHeaderTable()); |
| 682 | |
| 683 | // Create a Fetcher that uses our custom factory |
| 684 | auto fetcher = js.alloc<Fetcher>( |
| 685 | ioctx.addObject(kj::mv(outgoingFactory)), Fetcher::RequiresHostAndProtocol::YES); |
| 686 | |
| 687 | return kj::mv(fetcher); |
| 688 | })); |
| 689 | } |
| 690 | } // namespace workerd::api |