File
Blob: src/workerd/io/external-pusher.c++
| 1 | // Copyright (c) 2025 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 <workerd/io/external-pusher.h> |
| 6 | #include <workerd/jsg/jsg.h> |
| 7 | |
| 8 | namespace workerd { |
| 9 | |
| 10 | // ======================================================================================= |
| 11 | // ReadableStream handling |
| 12 | |
| 13 | namespace { |
| 14 | |
| 15 | // TODO(cleanup): These classes have been copied from streams/readable.c++. The copies there can be |
| 16 | // deleted as soon as we've switched from StreamSink to ExternalPusher and can delete all the |
| 17 | // StreamSink-related code. For now I'm not trying to avoid duplication. |
| 18 | |
| 19 | // HACK: We need as async pipe, like kj::newOneWayPipe(), except supporting explicit end(). So we |
| 20 | // wrap the two ends of the pipe in special adapters that track whether end() was called. |
| 21 | class ExplicitEndOutputPipeAdapter final: public capnp::ExplicitEndOutputStream { |
| 22 | public: |
| 23 | ExplicitEndOutputPipeAdapter( |
| 24 | kj::Own<kj::AsyncOutputStream> inner, kj::Own<kj::RefcountedWrapper<bool>> ended) |
| 25 | : inner(kj::mv(inner)), |
| 26 | ended(kj::mv(ended)) {} |
| 27 | |
| 28 | kj::Promise<void> write(kj::ArrayPtr<const byte> buffer) override { |
| 29 | return KJ_REQUIRE_NONNULL(inner)->write(buffer); |
| 30 | } |
| 31 | kj::Promise<void> write(kj::ArrayPtr<const kj::ArrayPtr<const byte>> pieces) override { |
| 32 | return KJ_REQUIRE_NONNULL(inner)->write(pieces); |
| 33 | } |
| 34 | |
| 35 | kj::Maybe<kj::Promise<uint64_t>> tryPumpFrom( |
| 36 | kj::AsyncInputStream& input, uint64_t amount) override { |
| 37 | return KJ_REQUIRE_NONNULL(inner)->tryPumpFrom(input, amount); |
| 38 | } |
| 39 | |
| 40 | kj::Promise<void> whenWriteDisconnected() override { |
| 41 | return KJ_REQUIRE_NONNULL(inner)->whenWriteDisconnected(); |
| 42 | } |
| 43 | |
| 44 | kj::Promise<void> end() override { |
| 45 | // Signal to the other side that end() was actually called. |
| 46 | ended->getWrapped() = true; |
| 47 | inner = kj::none; |
| 48 | return kj::READY_NOW; |
| 49 | } |
| 50 | |
| 51 | private: |
| 52 | kj::Maybe<kj::Own<kj::AsyncOutputStream>> inner; |
| 53 | kj::Own<kj::RefcountedWrapper<bool>> ended; |
| 54 | }; |
| 55 | |
| 56 | class ExplicitEndInputPipeAdapter final: public kj::AsyncInputStream { |
| 57 | public: |
| 58 | ExplicitEndInputPipeAdapter(kj::Own<kj::AsyncInputStream> inner, |
| 59 | kj::Own<kj::RefcountedWrapper<bool>> ended, |
| 60 | kj::Maybe<uint64_t> expectedLength) |
| 61 | : inner(kj::mv(inner)), |
| 62 | ended(kj::mv(ended)), |
| 63 | expectedLength(expectedLength) {} |
| 64 | |
| 65 | kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override { |
| 66 | size_t result = co_await inner->tryRead(buffer, minBytes, maxBytes); |
| 67 | |
| 68 | KJ_IF_SOME(l, expectedLength) { |
| 69 | KJ_ASSERT(result <= l); |
| 70 | l -= result; |
| 71 | if (l == 0) { |
| 72 | // If we got all the bytes we expected, we treat this as a successful end, because the |
| 73 | // underlying KJ pipe is not actually going to wait for the other side to drop. This is |
| 74 | // consistent with the behavior of Content-Length in HTTP anyway. |
| 75 | ended->getWrapped() = true; |
| 76 | } |
| 77 | } |
| 78 | |
| 79 | if (result < minBytes) { |
| 80 | // Verify that end() was called. |
| 81 | if (!ended->getWrapped()) { |
| 82 | JSG_FAIL_REQUIRE(Error, "ReadableStream received over RPC disconnected prematurely."); |
| 83 | } |
| 84 | } |
| 85 | co_return result; |
| 86 | } |
| 87 | |
| 88 | kj::Maybe<uint64_t> tryGetLength() override { |
| 89 | return inner->tryGetLength(); |
| 90 | } |
| 91 | |
| 92 | kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override { |
| 93 | return inner->pumpTo(output, amount); |
| 94 | } |
| 95 | |
| 96 | private: |
| 97 | kj::Own<kj::AsyncInputStream> inner; |
| 98 | kj::Own<kj::RefcountedWrapper<bool>> ended; |
| 99 | kj::Maybe<uint64_t> expectedLength; |
| 100 | }; |
| 101 | |
| 102 | } // namespace |
| 103 | |
| 104 | class ExternalPusherImpl::InputStreamImpl final: public ExternalPusher::InputStream::Server { |
| 105 | public: |
| 106 | InputStreamImpl(kj::Own<kj::AsyncInputStream> stream): stream(kj::mv(stream)) {} |
| 107 | |
| 108 | kj::Maybe<kj::Own<kj::AsyncInputStream>> stream; |
| 109 | }; |
| 110 | |
| 111 | kj::Promise<void> ExternalPusherImpl::pushByteStream(PushByteStreamContext context) { |
| 112 | kj::Maybe<uint64_t> expectedLength; |
| 113 | auto lp1 = context.getParams().getLengthPlusOne(); |
| 114 | if (lp1 > 0) { |
| 115 | expectedLength = lp1 - 1; |
| 116 | } |
| 117 | |
| 118 | auto pipe = kj::newOneWayPipe(expectedLength); |
| 119 | |
| 120 | auto endedFlag = kj::refcounted<kj::RefcountedWrapper<bool>>(false); |
| 121 | |
| 122 | auto out = kj::heap<ExplicitEndOutputPipeAdapter>(kj::mv(pipe.out), kj::addRef(*endedFlag)); |
| 123 | auto in = |
| 124 | kj::heap<ExplicitEndInputPipeAdapter>(kj::mv(pipe.in), kj::mv(endedFlag), expectedLength); |
| 125 | |
| 126 | auto results = context.initResults(capnp::MessageSize{4, 2}); |
| 127 | |
| 128 | results.setSource(inputStreamSet.add(kj::heap<InputStreamImpl>(kj::mv(in)))); |
| 129 | results.setSink(byteStreamFactory.kjToCapnp(kj::mv(out))); |
| 130 | return kj::READY_NOW; |
| 131 | } |
| 132 | |
| 133 | kj::Own<kj::AsyncInputStream> ExternalPusherImpl::unwrapStream( |
| 134 | ExternalPusher::InputStream::Client cap, kj::LiteralStringConst debugContext) { |
| 135 | return kj::newPromisedStream(unwrapStreamImpl(kj::mv(cap), debugContext)); |
| 136 | } |
| 137 | |
| 138 | kj::Promise<kj::Own<kj::AsyncInputStream>> ExternalPusherImpl::unwrapStreamImpl( |
| 139 | ExternalPusher::InputStream::Client cap, kj::LiteralStringConst debugContext) { |
| 140 | auto& unwrapped = KJ_REQUIRE_NONNULL(co_await inputStreamSet.getLocalServer(cap), |
| 141 | "pushed external is not a byte stream", debugContext, cap.debugInfo()); |
| 142 | |
| 143 | co_return KJ_REQUIRE_NONNULL(kj::mv(kj::downcast<InputStreamImpl>(unwrapped).stream), |
| 144 | "pushed byte stream has already been consumed"); |
| 145 | } |
| 146 | // ======================================================================================= |
| 147 | // AbortSignal handling |
| 148 | |
| 149 | namespace { |
| 150 | |
| 151 | // The jsrpc handler that receives aborts from the remote and triggers them locally |
| 152 | // |
| 153 | // TODO(cleanup): This class has been copied from basics.c++. The copy there can be |
| 154 | // deleted as soon as we've switched from StreamSink to ExternalPusher and can delete all the |
| 155 | // StreamSink-related code. For now I'm not trying to avoid duplication. |
| 156 | class AbortTriggerRpcServer final: public rpc::AbortTrigger::Server { |
| 157 | public: |
| 158 | AbortTriggerRpcServer(kj::Own<kj::PromiseFulfiller<void>> fulfiller, |
| 159 | kj::Own<ExternalPusherImpl::PendingAbortReason>&& pendingReason) |
| 160 | : fulfiller(kj::mv(fulfiller)), |
| 161 | pendingReason(kj::mv(pendingReason)) {} |
| 162 | |
| 163 | kj::Promise<void> abort(AbortContext abortCtx) override { |
| 164 | auto params = abortCtx.getParams(); |
| 165 | auto reason = params.getReason().getV8Serialized(); |
| 166 | |
| 167 | pendingReason->getWrapped() = kj::heapArray(reason.asBytes()); |
| 168 | fulfiller->fulfill(); |
| 169 | return kj::READY_NOW; |
| 170 | } |
| 171 | |
| 172 | kj::Promise<void> release(ReleaseContext releaseCtx) override { |
| 173 | released = true; |
| 174 | return kj::READY_NOW; |
| 175 | } |
| 176 | |
| 177 | ~AbortTriggerRpcServer() noexcept(false) { |
| 178 | if (pendingReason->getWrapped() != nullptr) { |
| 179 | // Already triggered |
| 180 | return; |
| 181 | } |
| 182 | |
| 183 | if (!released) { |
| 184 | pendingReason->getWrapped() = JSG_KJ_EXCEPTION(FAILED, DOMAbortError, |
| 185 | "An AbortSignal received over RPC was implicitly aborted because the connection back to " |
| 186 | "its trigger was lost."); |
| 187 | } |
| 188 | |
| 189 | // Always fulfill the promise in case the AbortSignal was waiting |
| 190 | fulfiller->fulfill(); |
| 191 | } |
| 192 | |
| 193 | private: |
| 194 | kj::Own<kj::PromiseFulfiller<void>> fulfiller; |
| 195 | kj::Own<ExternalPusherImpl::PendingAbortReason> pendingReason; |
| 196 | bool released = false; |
| 197 | }; |
| 198 | |
| 199 | } // namespace |
| 200 | |
| 201 | class ExternalPusherImpl::AbortSignalImpl final: public ExternalPusher::AbortSignal::Server { |
| 202 | public: |
| 203 | AbortSignalImpl(kj::Own<kj::PromiseFulfiller<rpc::AbortTrigger::Client>> triggerFulfiller) |
| 204 | : triggerFulfiller(kj::mv(triggerFulfiller)) {} |
| 205 | |
| 206 | kj::Maybe<kj::Own<kj::PromiseFulfiller<rpc::AbortTrigger::Client>>> triggerFulfiller; |
| 207 | }; |
| 208 | |
| 209 | kj::Promise<void> ExternalPusherImpl::pushAbortSignal(PushAbortSignalContext context) { |
| 210 | // We can't allocate the `AbortTriggerRpcServer` yet because we don't have the |
| 211 | // `PendingAbortReason` box, because that MUST be allocated inside unwrapAbortSignal() so it |
| 212 | // can be returned synchronously. So, we'll return a promise for a future `AbortTrigger`, and |
| 213 | // we'll put the fulfiller into the `AbortSignalImpl` where `unwrapAbortSignalImpl()` can find |
| 214 | // and fulfill it. |
| 215 | auto triggerPaf = kj::newPromiseAndFulfiller<rpc::AbortTrigger::Client>(); |
| 216 | |
| 217 | auto results = context.initResults(capnp::MessageSize{4, 2}); |
| 218 | results.setTrigger(kj::mv(triggerPaf.promise)); |
| 219 | results.setSignal(abortSignalSet.add(kj::heap<AbortSignalImpl>(kj::mv(triggerPaf.fulfiller)))); |
| 220 | |
| 221 | return kj::READY_NOW; |
| 222 | } |
| 223 | |
| 224 | ExternalPusherImpl::AbortSignal ExternalPusherImpl::unwrapAbortSignal( |
| 225 | ExternalPusher::AbortSignal::Client cap) { |
| 226 | // We need to return a result synchronously, including the PendingAbortReason box. But, |
| 227 | // pushAbortSignal() might not have been received yet. So, we have to allocate the box here, so |
| 228 | // we can return it. Then we can try to wire it up to the right trigger later, in |
| 229 | // unwrapAbortSignalImpl(). |
| 230 | auto pendingReason = kj::refcounted<PendingAbortReason>(); |
| 231 | auto promise = unwrapAbortSignalImpl(kj::mv(cap), kj::addRef(*pendingReason)); |
| 232 | |
| 233 | return { |
| 234 | .signal = kj::mv(promise), |
| 235 | .reason = kj::mv(pendingReason), |
| 236 | }; |
| 237 | } |
| 238 | |
| 239 | kj::Promise<void> ExternalPusherImpl::unwrapAbortSignalImpl( |
| 240 | ExternalPusher::AbortSignal::Client cap, kj::Own<PendingAbortReason> pendingReason) { |
| 241 | auto paf = kj::newPromiseAndFulfiller<void>(); |
| 242 | |
| 243 | { |
| 244 | auto& unwrapped = KJ_REQUIRE_NONNULL( |
| 245 | co_await abortSignalSet.getLocalServer(cap), "pushed external is not an AbortSignal"); |
| 246 | |
| 247 | auto triggerFulfiller = |
| 248 | KJ_REQUIRE_NONNULL(kj::mv(kj::downcast<AbortSignalImpl>(unwrapped).triggerFulfiller), |
| 249 | "pushed AbortSignal has already been consumed"); |
| 250 | |
| 251 | triggerFulfiller->fulfill( |
| 252 | kj::heap<AbortTriggerRpcServer>(kj::mv(paf.fulfiller), kj::mv(pendingReason))); |
| 253 | |
| 254 | // We don't need `cap` anymore. |
| 255 | auto drop = kj::mv(cap); |
| 256 | } |
| 257 | |
| 258 | co_await paf.promise; |
| 259 | } |
| 260 | |
| 261 | } // namespace workerd |