Skip to content

Commit afc70fa

Browse files
Defer TCP socket connect() until DO output gate clears
Adds TCP_SOCKET_CONNECT_OUTPUT_GATE autogate that defers httpClient->connect() until IoContext::waitForOutputLocks() resolves. Socket is returned immediately backed by kj::newPromisedStream so reads/writes queue until the real connection arrives.
1 parent 1e952a5 commit afc70fa

8 files changed

Lines changed: 205 additions & 20 deletions

File tree

src/workerd/api/sockets.c++

Lines changed: 91 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
#include <workerd/io/worker-interface.h>
1212
#include <workerd/jsg/exception.h>
1313
#include <workerd/jsg/url.h>
14+
#include <workerd/util/autogate.h>
1415

1516
namespace workerd::api {
1617

@@ -90,7 +91,8 @@ jsg::Ref<Socket> setupSocket(jsg::Lock& js,
9091
SecureTransportKind secureTransport,
9192
kj::Maybe<kj::String> domain,
9293
bool isDefaultFetchPort,
93-
kj::Maybe<jsg::PromiseResolverPair<SocketInfo>> maybeOpenedPrPair) {
94+
kj::Maybe<jsg::PromiseResolverPair<SocketInfo>> maybeOpenedPrPair,
95+
kj::Maybe<kj::Promise<void>> connectTask) {
9496
auto& ioContext = IoContext::current();
9597

9698
// Disconnection handling is annoyingly complicated:
@@ -173,7 +175,7 @@ jsg::Ref<Socket> setupSocket(jsg::Lock& js,
173175
auto result = js.alloc<Socket>(js, ioContext, kj::mv(refcountedConnection), kj::mv(remoteAddress),
174176
kj::mv(readable), kj::mv(writable), kj::mv(closedPrPair), kj::mv(watchForDisconnectTask),
175177
kj::mv(options), kj::mv(tlsStarter), secureTransport, kj::mv(domain), isDefaultFetchPort,
176-
kj::mv(openedPrPair));
178+
kj::mv(openedPrPair), kj::mv(connectTask));
177179

178180
KJ_IF_SOME(p, eofPromise) {
179181
result->handleReadableEof(js, kj::mv(p));
@@ -256,25 +258,99 @@ jsg::Ref<Socket> connectImplNoOutputLock(jsg::Lock& js,
256258
}
257259
kj::Own<kj::TlsStarterCallback> tlsStarter = kj::heap<kj::TlsStarterCallback>();
258260
httpConnectSettings.tlsStarter = tlsStarter;
259-
auto request = httpClient->connect(addressStr, *headers, httpConnectSettings);
260-
request.connection = request.connection.attach(kj::mv(httpClient));
261-
262-
auto result = setupSocket(js, kj::mv(request.connection), kj::mv(addressStr), kj::mv(options),
263-
kj::mv(tlsStarter), secureTransport, kj::mv(domain), isDefaultFetchPort,
264-
kj::none /* maybeOpenedPrPair */);
265-
// `handleProxyStatus` needs an initialized refcount to use `JSG_THIS`, hence it cannot be
266-
// called in Socket's constructor. Also it's only necessary when creating a Socket as a result of
267-
// a `connect`.
268-
result->handleProxyStatus(js, kj::mv(request.status));
269-
return result;
261+
262+
if (util::Autogate::isEnabled(util::AutogateKey::TCP_SOCKET_CONNECT_OUTPUT_GATE)) {
263+
// Deferred-connect path: return a Socket backed by a promised stream immediately, deferring
264+
// the actual httpClient->connect() until after the DO output gate clears. This prevents
265+
// premature network output while storage writes are still pending.
266+
//
267+
// Two promise-fulfiller pairs bridge the deferred connect to the Socket:
268+
// connStreamPaf — fulfilled with the raw connection stream once connect() returns
269+
// statusPaf — fulfilled with the proxy status once connect()'s status promise resolves
270+
//
271+
// The connect task is stored inside ConnectionData alongside tlsStarter. This ensures
272+
// that if the Socket is GC'd (destroying ConnectionData), the connect task is cancelled
273+
// before tlsStarter is destroyed, preventing a use-after-free on the TlsStarterCallback
274+
// pointer embedded in httpConnectSettings.
275+
276+
// Promise-fulfiller for the raw connection stream.
277+
auto connStreamPaf = kj::newPromiseAndFulfiller<kj::Own<kj::AsyncIoStream>>();
278+
auto& connStreamFulfiller = *connStreamPaf.fulfiller;
279+
auto deferredCancelStream = kj::defer([fulfiller = kj::mv(connStreamPaf.fulfiller)]() mutable {
280+
fulfiller->reject(KJ_EXCEPTION(DISCONNECTED, "socket connect cancelled"));
281+
});
282+
283+
// Promise-fulfiller for the proxy status.
284+
auto statusPaf = kj::newPromiseAndFulfiller<kj::HttpClient::ConnectRequest::Status>();
285+
auto& statusFulfiller = *statusPaf.fulfiller;
286+
auto deferredCancelStatus = kj::defer([fulfiller = kj::mv(statusPaf.fulfiller)]() mutable {
287+
fulfiller->reject(KJ_EXCEPTION(DISCONNECTED, "socket connect cancelled"));
288+
});
289+
290+
// The connect task: waits for the output gate, performs the connect, then forwards results.
291+
static auto constexpr doConnect =
292+
[](IoContext& ioContext, kj::StringPtr addressStr, kj::HttpHeaders& headers,
293+
kj::HttpConnectSettings httpConnectSettings, kj::Own<kj::HttpClient> httpClient,
294+
kj::PromiseFulfiller<kj::Own<kj::AsyncIoStream>>& connFulfiller,
295+
kj::PromiseFulfiller<kj::HttpClient::ConnectRequest::Status>& statusFulfiller)
296+
-> kj::Promise<void> {
297+
try {
298+
co_await ioContext.waitForOutputLocks();
299+
auto request = httpClient->connect(addressStr, headers, httpConnectSettings);
300+
request.connection = request.connection.attach(kj::mv(httpClient));
301+
auto status = kj::mv(request.status);
302+
connFulfiller.fulfill(kj::mv(request.connection));
303+
auto resolvedStatus = co_await kj::mv(status);
304+
statusFulfiller.fulfill(kj::mv(resolvedStatus));
305+
} catch (...) {
306+
auto e = kj::getCaughtExceptionAsKj();
307+
if (!connFulfiller.isWaiting()) {
308+
// Connection was already fulfilled; only reject status.
309+
statusFulfiller.reject(kj::mv(e));
310+
} else {
311+
connFulfiller.reject(kj::cp(e));
312+
statusFulfiller.reject(kj::mv(e));
313+
}
314+
}
315+
};
316+
317+
auto connectTaskPromise =
318+
doConnect(ioContext, addressStr, *headers, httpConnectSettings, kj::mv(httpClient),
319+
connStreamFulfiller, statusFulfiller)
320+
.attach(kj::mv(deferredCancelStream), kj::mv(deferredCancelStatus), kj::mv(headers));
321+
322+
auto result = setupSocket(js, kj::newPromisedStream(kj::mv(connStreamPaf.promise)),
323+
kj::mv(addressStr), kj::mv(options), kj::mv(tlsStarter), secureTransport, kj::mv(domain),
324+
isDefaultFetchPort, kj::none /* maybeOpenedPrPair */, kj::mv(connectTaskPromise));
325+
326+
// `handleProxyStatus` needs an initialized refcount to use `JSG_THIS`, hence it cannot be
327+
// called in Socket's constructor. Also it's only necessary when creating a Socket as a result
328+
// of a `connect`.
329+
result->handleProxyStatus(js, kj::mv(statusPaf.promise));
330+
return result;
331+
} else {
332+
// Original synchronous-connect path (no output gate wait).
333+
auto request = httpClient->connect(addressStr, *headers, httpConnectSettings);
334+
request.connection = request.connection.attach(kj::mv(httpClient));
335+
336+
auto result = setupSocket(js, kj::mv(request.connection), kj::mv(addressStr), kj::mv(options),
337+
kj::mv(tlsStarter), secureTransport, kj::mv(domain), isDefaultFetchPort,
338+
kj::none /* maybeOpenedPrPair */);
339+
// `handleProxyStatus` needs an initialized refcount to use `JSG_THIS`, hence it cannot be
340+
// called in Socket's constructor. Also it's only necessary when creating a Socket as a result
341+
// of a `connect`.
342+
result->handleProxyStatus(js, kj::mv(request.status));
343+
return result;
344+
}
270345
}
271346

272347
jsg::Ref<Socket> connectImpl(jsg::Lock& js,
273348
kj::Maybe<jsg::Ref<Fetcher>> fetcher,
274349
AnySocketAddress address,
275350
jsg::Optional<SocketOptions> options) {
276-
// TODO(soon): Doesn't this need to check for the presence of an output lock, and if it finds one
277-
// then wait on it, before calling into connectImplNoOutputLock?
351+
// When the TCP_SOCKET_CONNECT_OUTPUT_GATE autogate is enabled, the output gate wait is
352+
// handled inside connectImplNoOutputLock via a deferred connect task, so no separate wait
353+
// is needed here.
278354
return connectImplNoOutputLock(js, kj::mv(fetcher), kj::mv(address), kj::mv(options));
279355
}
280356

src/workerd/api/sockets.h

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -70,9 +70,12 @@ class Socket: public jsg::Object {
7070
SecureTransportKind secureTransport,
7171
kj::Maybe<kj::String> domain,
7272
bool isDefaultFetchPort,
73-
jsg::PromiseResolverPair<SocketInfo> openedPrPair)
74-
: connectionData(context.addObject(kj::heap<ConnectionData>(
75-
kj::mv(tlsStarter), kj::mv(connectionStream), kj::mv(watchForDisconnectTask)))),
73+
jsg::PromiseResolverPair<SocketInfo> openedPrPair,
74+
kj::Maybe<kj::Promise<void>> connectTask = kj::none)
75+
: connectionData(context.addObject(kj::heap<ConnectionData>(kj::mv(tlsStarter),
76+
kj::mv(connectionStream),
77+
kj::mv(watchForDisconnectTask),
78+
kj::mv(connectTask)))),
7679
readable(kj::mv(readableParam)),
7780
writable(kj::mv(writable)),
7881
closedResolver(kj::mv(closedPrPair.resolver)),
@@ -179,14 +182,21 @@ class Socket: public jsg::Object {
179182
struct ConnectionData {
180183
kj::Own<kj::RefcountedWrapper<kj::Own<kj::AsyncIoStream>>> connectionStream;
181184
kj::Maybe<kj::Promise<void>> watchForDisconnectTask;
185+
// When the deferred-connect autogate is enabled, holds the task that waits for the
186+
// output gate then calls httpClient->connect(). Lives here so that cancelling the
187+
// Socket (destroying ConnectionData) also cancels the connect task, preventing a
188+
// TlsStarterCallback use-after-free.
189+
kj::Maybe<kj::Promise<void>> connectTask;
182190
// tlsStarter must be declared after connectionStream so that it is destroyed first,
183191
// since it holds a reference that keeps the connection alive.
184192
kj::Own<kj::TlsStarterCallback> tlsStarter;
185193
ConnectionData(kj::Own<kj::TlsStarterCallback> tlsStarter,
186194
kj::Own<kj::RefcountedWrapper<kj::Own<kj::AsyncIoStream>>> connStream,
187-
kj::Promise<void> disconnectTask)
195+
kj::Promise<void> disconnectTask,
196+
kj::Maybe<kj::Promise<void>> connectTask = kj::none)
188197
: connectionStream(kj::mv(connStream)),
189198
watchForDisconnectTask(kj::mv(disconnectTask)),
199+
connectTask(kj::mv(connectTask)),
190200
tlsStarter(kj::mv(tlsStarter)) {}
191201
};
192202
kj::Maybe<IoOwn<ConnectionData>> connectionData;
@@ -251,7 +261,8 @@ jsg::Ref<Socket> setupSocket(jsg::Lock& js,
251261
SecureTransportKind secureTransport,
252262
kj::Maybe<kj::String> domain,
253263
bool isDefaultFetchPort,
254-
kj::Maybe<jsg::PromiseResolverPair<SocketInfo>> maybeOpenedPrPair);
264+
kj::Maybe<jsg::PromiseResolverPair<SocketInfo>> maybeOpenedPrPair,
265+
kj::Maybe<kj::Promise<void>> connectTask = kj::none);
255266

256267
jsg::Ref<Socket> connectImplNoOutputLock(jsg::Lock& js,
257268
kj::Maybe<jsg::Ref<Fetcher>> fetcher,

src/workerd/api/tests/BUILD.bazel

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -630,6 +630,19 @@ wd_test(
630630
],
631631
)
632632

633+
js_binary(
634+
name = "socket-output-gate-server",
635+
entry_point = "socket-output-gate-server.js",
636+
)
637+
638+
wd_test(
639+
src = "socket-output-gate-test.wd-test",
640+
args = ["--experimental"],
641+
data = ["socket-output-gate-test.js"],
642+
sidecar = "socket-output-gate-server",
643+
sidecar_port_bindings = ["ECHO_SERVER_PORT"],
644+
)
645+
633646
wd_test(
634647
src = "js-rpc-socket-test.wd-test",
635648
args = [
Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
// Simple TCP echo server for socket-output-gate test.
2+
const net = require('node:net');
3+
4+
const server = net.createServer((socket) => {
5+
socket.on('data', (chunk) => socket.write(chunk));
6+
socket.on('error', () => {});
7+
});
8+
9+
server.listen(process.env.ECHO_SERVER_PORT, () => {
10+
console.info(`Echo server on port ${server.address().port}`);
11+
});
Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
// Tests that TCP socket connect() respects the DO output gate when
2+
// TCP_SOCKET_CONNECT_OUTPUT_GATE is enabled.
3+
//
4+
// Invariant: by the time socket.opened resolves, any storage write started
5+
// in the same turn must already be committed.
6+
7+
import { connect } from 'cloudflare:sockets';
8+
import { strict as assert } from 'node:assert';
9+
import { DurableObject } from 'cloudflare:workers';
10+
11+
export class SocketTestDO extends DurableObject {
12+
async fetch(request) {
13+
const port = new URL(request.url).searchParams.get('port');
14+
15+
let putCompleted = false;
16+
17+
// Storage write — no await. This creates a pending output gate lock.
18+
this.ctx.storage.put('key', 'value').then(() => {
19+
putCompleted = true;
20+
});
21+
22+
// Connect immediately while the write may still be in-flight.
23+
const socket = connect(`localhost:${port}`);
24+
25+
// socket.opened only resolves after handleProxyStatus confirms a successful
26+
// proxy status — which, with the autogate, cannot happen until
27+
// waitForOutputLocks() resolves, i.e. after the put is committed.
28+
await socket.opened;
29+
30+
assert.ok(putCompleted, 'storage put must complete before socket opens');
31+
32+
await socket.close();
33+
return new Response('ok');
34+
}
35+
}
36+
37+
export const connectRespectsOutputGate = {
38+
async test(ctrl, env) {
39+
const id = env.SOCKET_TEST_DO.newUniqueId();
40+
const stub = env.SOCKET_TEST_DO.get(id);
41+
const resp = await stub.fetch(`http://do/?port=${env.ECHO_SERVER_PORT}`);
42+
assert.equal(resp.status, 200);
43+
},
44+
};
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
using Workerd = import "/workerd/workerd.capnp";
2+
3+
const unitTests :Workerd.Config = (
4+
autogates = ["workerd-autogate-tcp-socket-connect-output-gate"],
5+
services = [
6+
(
7+
name = "main",
8+
worker = (
9+
modules = [
10+
(name = "worker", esModule = embed "socket-output-gate-test.js"),
11+
],
12+
compatibilityFlags = ["nodejs_compat", "experimental"],
13+
bindings = [
14+
(name = "ECHO_SERVER_PORT", fromEnvironment = "ECHO_SERVER_PORT"),
15+
(name = "SOCKET_TEST_DO", durableObjectNamespace = "SocketTestDO"),
16+
],
17+
durableObjectNamespaces = [
18+
(className = "SocketTestDO", uniqueKey = "210bd0cbd803ef7883a1ee9d86cce06e"),
19+
],
20+
durableObjectStorage = (inMemory = void),
21+
)
22+
),
23+
( name = "internet", network = ( allow = ["private"] ) ),
24+
],
25+
);

src/workerd/util/autogate.c++

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,8 @@ kj::StringPtr KJ_STRINGIFY(AutogateKey key) {
3535
return "enable-fast-textencoder"_kj;
3636
case AutogateKey::ENABLE_DRAINING_READ_ON_STANDARD_STREAMS:
3737
return "enable-draining-read-on-standard-streams"_kj;
38+
case AutogateKey::TCP_SOCKET_CONNECT_OUTPUT_GATE:
39+
return "tcp-socket-connect-output-gate"_kj;
3840
case AutogateKey::NumOfKeys:
3941
KJ_FAIL_ASSERT("NumOfKeys should not be used in getName");
4042
}

src/workerd/util/autogate.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,9 @@ enum class AutogateKey {
3939
ENABLE_FAST_TEXTENCODER,
4040
// Enable draining read on standard streams
4141
ENABLE_DRAINING_READ_ON_STANDARD_STREAMS,
42+
// Defers TCP socket connect() to wait for DO output gate, preventing
43+
// network outputs while storage writes are pending.
44+
TCP_SOCKET_CONNECT_OUTPUT_GATE,
4245
NumOfKeys // Reserved for iteration.
4346
};
4447

0 commit comments

Comments
 (0)