diff --git a/src/workerd/api/sockets.c++ b/src/workerd/api/sockets.c++ index 876ef59bd78..6d45532d19d 100644 --- a/src/workerd/api/sockets.c++ +++ b/src/workerd/api/sockets.c++ @@ -21,6 +21,7 @@ #include #include +#include namespace workerd::api { @@ -351,6 +352,154 @@ JsWritableStream newDatagramWritableStream(jsg::Lock& js, IoOwn }); } +class RpcDatagramChannel final: public DatagramChannel, public kj::Refcounted { + public: + // Creates a channel that sends datagrams through an RPC stream. + RpcDatagramChannel(rpc::DatagramStream::Client down): down(kj::mv(down)) {} + + // Disconnects any outstanding channel operations. + ~RpcDatagramChannel() noexcept(false) { + disconnect(KJ_EXCEPTION(DISCONNECTED, "UDP RPC channel was destroyed")); + } + + // Receives the next datagram or end-of-stream. + kj::Promise>> receive() override { + KJ_IF_SOME(exception, failure) { + return exception.clone(); + } + KJ_IF_SOME(item, pending) { + auto datagram = kj::mv(item.datagram); + item.consumed->fulfill(); + pending = kj::none; + return kj::Maybe>(kj::mv(datagram)); + } + if (ended) { + return kj::Maybe>(kj::none); + } + KJ_REQUIRE(receivers.empty(), "DatagramChannel::receive() already has a pending call"); + return receivers.wait(); + } + + // Sends one datagram to the remote endpoint. + kj::Promise send(kj::ArrayPtr datagram) override { + KJ_IF_SOME(exception, failure) { + return exception.clone(); + } + auto req = down.sendRequest(capnp::MessageSize{8 + datagram.size() / sizeof(capnp::word), 0}); + req.setDatagram(datagram); + return req.send(); + } + + // Delivers one incoming RPC datagram with consumption backpressure. + kj::Promise deliver(kj::Array datagram) { + KJ_IF_SOME(exception, failure) { + return exception.clone(); + } + KJ_REQUIRE(!ended, "datagram received after UDP RPC stream ended"); + if (!receivers.empty()) { + receivers.fulfill(kj::Maybe>(kj::mv(datagram))); + return kj::READY_NOW; + } + + KJ_REQUIRE(pending == kj::none, "UDP RPC stream delivered concurrent datagrams"); + auto paf = kj::newPromiseAndFulfiller(); + pending = Pending{kj::mv(datagram), kj::mv(paf.fulfiller)}; + return kj::mv(paf.promise); + } + + // Marks the incoming datagram stream complete. + void endIncoming() { + KJ_REQUIRE(pending == kj::none, "UDP RPC stream ended with an undelivered datagram"); + ended = true; + if (!receivers.empty()) { + receivers.fulfill(kj::Maybe>(kj::none)); + } + } + + // Signals completion of the outgoing datagram stream. + kj::Promise endOutgoing() { + return down.endRequest().sendIgnoringResult(); + } + + // Fails outstanding operations and closes the RPC stream. + void disconnect(kj::Exception exception) { + if (failure != kj::none) return; + failure = exception.clone(); + if (!receivers.empty()) { + receivers.reject(exception.clone()); + } + KJ_IF_SOME(item, pending) { + item.consumed->reject(exception.clone()); + pending = kj::none; + } + down = nullptr; + } + + private: + struct Pending { + kj::Array datagram; + kj::Own> consumed; + }; + + rpc::DatagramStream::Client down; + kj::WaiterQueue>> receivers; + kj::Maybe pending; + kj::Maybe failure; + bool ended = false; +}; + +class IncomingRpcDatagramStream final: public rpc::DatagramStream::Server { + public: + IncomingRpcDatagramStream(kj::Rc channel): channel(kj::mv(channel)) {} + + private: + kj::Promise send(SendContext context) override { + auto datagram = context.getParams().getDatagram(); + return channel->deliver(kj::heapArray(datagram)); + } + + kj::Promise end(EndContext) override { + channel->endIncoming(); + return kj::READY_NOW; + } + + kj::Rc channel; +}; + +class OutgoingRpcDatagramStream final: public rpc::DatagramStream::Server { + public: + OutgoingRpcDatagramStream(kj::Rc channel): channel(kj::mv(channel)) {} + + private: + kj::Promise send(SendContext context) override { + KJ_REQUIRE(!ended, "datagram received after UDP RPC stream ended"); + return channel->send(context.getParams().getDatagram()); + } + + kj::Promise end(EndContext) override { + ended = true; + return kj::READY_NOW; + } + + kj::Rc channel; + bool ended = false; +}; + +kj::Promise pumpDatagramsToRpc( + kj::Rc channel, rpc::DatagramStream::Client stream) { + for (;;) { + KJ_IF_SOME(datagram, co_await channel->receive()) { + auto req = + stream.sendRequest(capnp::MessageSize{8 + datagram.size() / sizeof(capnp::word), 0}); + req.setDatagram(datagram); + co_await req.send(); + } else { + co_await stream.endRequest().sendIgnoringResult(); + co_return; + } + } +} + } // namespace // Forward declarations @@ -498,6 +647,48 @@ kj::Promise UdpConnectCustomEvent::run( co_return Result{.outcome = outcome}; } +kj::Promise UdpConnectCustomEvent::sendRpc( + capnp::HttpOverCapnpFactory&, + capnp::ByteStreamFactory&, + FrankenvalueHandler&, + rpc::EventDispatcher::Client dispatcher) { + auto rpcChannel = newNeuterableDatagramChannel(channel); + KJ_DEFER(rpcChannel->neuter(KJ_EXCEPTION(DISCONNECTED, "UDP RPC event ended"))); + + auto req = dispatcher.udpConnectRequest(); + req.setHost(host); + req.setDown(kj::heap(rpcChannel.addRef())); + auto sent = req.send(); + auto up = sent.getUp(); + + EventOutcome outcome = EventOutcome::UNKNOWN; + auto responseTask = sent.then([&outcome](auto response) { outcome = response.getResult(); }); + auto pumpTask = + pumpDatagramsToRpc(rpcChannel.addRef(), kj::mv(up)).then([]() -> kj::Promise { + return kj::NEVER_DONE; + }); + co_await responseTask.exclusiveJoin(kj::mv(pumpTask)); + co_return Result{.outcome = outcome}; +} + +kj::Promise UdpConnectCustomEvent::receiveRpc( + UdpConnectContext context, WorkerInterface& worker) { + auto params = context.getParams(); + auto channel = kj::rc(params.getDown()); + KJ_DEFER(channel->disconnect(KJ_EXCEPTION(DISCONNECTED, "UDP RPC event ended"))); + + rpc::DatagramStream::Client up = kj::heap(channel.addRef()); + capnp::PipelineBuilder pipelineBuilder; + pipelineBuilder.setUp(kj::cp(up)); + context.setPipeline(pipelineBuilder.build()); + context.getResults(capnp::MessageSize{4, 1}).setUp(kj::mv(up)); + + auto event = kj::heap(kj::str(params.getHost()), *channel); + auto result = co_await worker.customEvent(kj::mv(event)); + co_await channel->endOutgoing(); + context.getResults().setResult(result.outcome); +} + tracing::EventInfo UdpConnectCustomEvent::getEventInfo() const { return tracing::ConnectEventInfo(); } diff --git a/src/workerd/api/sockets.h b/src/workerd/api/sockets.h index 3a455f40c75..bc55bd478b4 100644 --- a/src/workerd/api/sockets.h +++ b/src/workerd/api/sockets.h @@ -430,10 +430,6 @@ jsg::Ref setupDatagramSocket(jsg::Lock& js, // WorkerInterface::customEvent() with it, exactly as Queue/Alarm/Scheduled events do for their own // non-HTTP-shaped triggers. // -// This event cannot be forwarded over RPC: a DatagramChannel is a live, in-process-only object, -// so sendRpc() is unimplemented. It is only ever dispatched by a listener running in the same -// process as the worker. -// // `channel` is borrowed, not owned: the listener that constructs this event attaches the // underlying flow's ownership to the same task that dispatches this event (see // Server::UdpListener::dispatch()), so it is guaranteed to outlive every call made through this @@ -454,11 +450,11 @@ class UdpConnectCustomEvent final: public WorkerInterface::CustomEvent { kj::Promise sendRpc(capnp::HttpOverCapnpFactory& httpOverCapnpFactory, capnp::ByteStreamFactory& byteStreamFactory, FrankenvalueHandler& frankenvalueHandler, - rpc::EventDispatcher::Client dispatcher) override { - KJ_UNIMPLEMENTED( - "a UDP connect event cannot be forwarded over RPC; it is only ever dispatched in-process " - "by the listener that owns the underlying datagram flow"); - } + rpc::EventDispatcher::Client dispatcher) override; + + using UdpConnectContext = capnp::CallContext; + static kj::Promise receiveRpc(UdpConnectContext context, WorkerInterface& worker); kj::Promise notSupported() override { KJ_UNIMPLEMENTED("udp connect event not supported"); diff --git a/src/workerd/io/worker-interface.capnp b/src/workerd/io/worker-interface.capnp index 93e52542e9d..87b540ab28b 100644 --- a/src/workerd/io/worker-interface.capnp +++ b/src/workerd/io/worker-interface.capnp @@ -878,6 +878,14 @@ interface TailStreamTarget $Cxx.allowCancellation { # Report one or more streaming tail events to a tail worker. } +interface DatagramStream $Cxx.allowCancellation { + send @0 (datagram :Data) -> stream; + # Sends one datagram. Each call preserves a message boundary. + + end @1 (); + # Signals that no more datagrams will be sent and reports errors from previous send() calls. +} + interface EventDispatcher @0xf20697475ec1752d { # Interface used to deliver events to a Worker's global event handlers. @@ -968,6 +976,12 @@ interface EventDispatcher @0xf20697475ec1752d { # instantiated there -- or maybe some mechanism for running "remote facets". For now, though, # we punt and simply don't support it.) + udpConnect @14 (host :Text, down :DatagramStream) + -> (up :DatagramStream, result :EventOutcome) $Cxx.allowCancellation; + # Opens a UDP flow. `up` carries datagrams received from the peer toward the Worker, while `down` + # carries datagrams sent by the Worker back toward the peer. The call remains pending until the + # Worker's connect() handler completes. + # Other methods might be added to handle other kinds of events, e.g. TCP connections, or maybe # even native Cap'n Proto RPC eventually. } diff --git a/src/workerd/server/server-test.c++ b/src/workerd/server/server-test.c++ index ca622de6e77..2590947218b 100644 --- a/src/workerd/server/server-test.c++ +++ b/src/workerd/server/server-test.c++ @@ -335,9 +335,14 @@ class TestServer final: private kj::Filesystem, private kj::EntropySource, priva kj::Own source; }; + struct SentDatagram { + kj::Array content; + kj::String destination; + }; + struct DatagramState final: public kj::Refcounted { kj::ProducerConsumerQueue incoming; - kj::ProducerConsumerQueue> outgoing; + kj::ProducerConsumerQueue outgoing; }; TestServer(kj::StringPtr configText, @@ -427,6 +432,8 @@ class TestServer final: private kj::Filesystem, private kj::EntropySource, priva bool hasUdp(kj::StringPtr addr); + SentDatagram receiveUdp(kj::StringPtr addr); + // Try to connect to the address and return whether or not this connection attempt hangs, // i.e. a listener exists but connections are not being accepted. bool connectHangs(kj::StringPtr addr) { @@ -544,7 +551,7 @@ class TestServer final: private kj::Filesystem, private kj::EntropySource, priva kj::Promise send( kj::ArrayPtr buffer, kj::NetworkAddress& destination) override { - state->outgoing.push(kj::heapArray(buffer)); + state->outgoing.push({kj::heapArray(buffer), destination.toString()}); return buffer.size(); } @@ -700,6 +707,12 @@ bool TestServer::hasUdp(kj::StringPtr addr) { return getDatagramState(addr)->outgoing.pop().poll(ws); } +TestServer::SentDatagram TestServer::receiveUdp(kj::StringPtr addr) { + auto datagram = getDatagramState(addr)->outgoing.pop(); + KJ_REQUIRE(datagram.poll(ws), "No UDP datagram available"); + return datagram.wait(ws); +} + // ======================================================================================= // Test Workers @@ -5585,6 +5598,43 @@ KJ_TEST("Server: JS RPC over HTTP connections") { conn.httpGet200("/", "got: 35"); } +KJ_TEST("Server: UDP RPC over HTTP connections") { + TestServer test(R"(( + services = [ + ( name = "worker", + worker = ( + compatibilityDate = "2024-02-23", + compatibilityFlags = ["experimental"], + modules = [( + name = "main.js", + esModule = + `export default { + ` async connect(socket) { + ` const { value } = await socket.readable.getReader().read(); + ` await socket.writable.getWriter().write(value); + ` } + `} + )] + ) + ), + (name = "outbound", external = (address = "loopback", http = (capnpConnectHost = "cappy"))) + ], + sockets = [ + ( name = "rpc", address = "loopback", service = "worker", + http = (capnpConnectHost = "cappy")), + ( name = "udp", address = "udp-address", service = "outbound", udp = ()), + ] + ))"_kj); + + test.server.allowExperimental(); + test.start(); + + test.sendUdp("udp-address", "peer:1234", "hello"_kjb); + auto response = test.receiveUdp("udp-address"); + KJ_EXPECT(response.content.asPtr() == "hello"_kjb); + KJ_EXPECT(response.destination == "peer:1234"); +} + KJ_TEST("Server: Entrypoint binding with props") { TestServer test(R"(( services = [ diff --git a/src/workerd/server/server.c++ b/src/workerd/server/server.c++ index 9da50a9891c..744076295ea 100644 --- a/src/workerd/server/server.c++ +++ b/src/workerd/server/server.c++ @@ -6435,6 +6435,12 @@ class Server::WorkerdBootstrapImpl final: public rpc::WorkerdBootstrap::Server { return api::JsRpcSessionCustomEvent::receiveRpc(context, getWorker()); } + kj::Promise udpConnect(UdpConnectContext context) override { + auto worker = getWorker(); + auto& workerRef = *worker; + return api::UdpConnectCustomEvent::receiveRpc(context, workerRef).attach(kj::mv(worker)); + } + kj::Promise tailStreamSession(TailStreamSessionContext context) override { auto customEvent = kj::heap(); auto cap = customEvent->getCap();