Skip to content
File

Blob: src/workerd/api/sockets.c++

30.5 KB
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 
17namespace workerd::api {
18 
19namespace {
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.
24bool 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 
52SecureTransportKind 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 
66bool 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 
75kj::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
85class StreamWorkerInterface;
86 
87jsg::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 
187jsg::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 
285jsg::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 
295jsg::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 
335jsg::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 
435void 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 
480void 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 
509void 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 
516void 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 
527jsg::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 
555jsg::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 
560kj::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
583class 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
601class 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 
657kj::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 
667jsg::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