File
Blob: src/workerd/api/streams/writable-sink.c++
| 1 | #include "writable-sink.h" |
| 2 | |
| 3 | #include <workerd/io/io-context.h> |
| 4 | #include <workerd/util/state-machine.h> |
| 5 | #include <workerd/util/stream-utils.h> |
| 6 | |
| 7 | #include <capnp/compat/byte-stream.h> |
| 8 | #include <kj/async-io.h> |
| 9 | #include <kj/compat/brotli.h> |
| 10 | #include <kj/compat/gzip.h> |
| 11 | |
| 12 | namespace workerd::api::streams { |
| 13 | |
| 14 | namespace { |
| 15 | struct Closed { |
| 16 | static constexpr kj::StringPtr NAME KJ_UNUSED = "closed"_kj; |
| 17 | }; |
| 18 | |
| 19 | struct Open { |
| 20 | static constexpr kj::StringPtr NAME KJ_UNUSED = "open"_kj; |
| 21 | kj::Own<kj::AsyncOutputStream> stream; |
| 22 | }; |
| 23 | |
| 24 | // State machine for tracking writable sink lifecycle: |
| 25 | // Open -> Closed (normal close via end()) |
| 26 | // Open -> kj::Exception (error via abort() or write failure) |
| 27 | // Closed is terminal, kj::Exception is implicitly terminal via ErrorState. |
| 28 | using WritableSinkState = StateMachine<TerminalStates<Closed>, |
| 29 | ErrorState<kj::Exception>, |
| 30 | ActiveState<Open>, |
| 31 | Open, |
| 32 | Closed, |
| 33 | kj::Exception>; |
| 34 | |
| 35 | // The base implementation of WritableSink. This is not exposed publicly. |
| 36 | class WritableSinkImpl: public WritableSink { |
| 37 | public: |
| 38 | WritableSinkImpl(kj::Own<kj::AsyncOutputStream> inner, |
| 39 | rpc::StreamEncoding encoding = rpc::StreamEncoding::IDENTITY) |
| 40 | : state(WritableSinkState::create<Open>(kj::mv(inner))), |
| 41 | encoding(encoding) {} |
| 42 | WritableSinkImpl() |
| 43 | : state(WritableSinkState::create<Closed>()), |
| 44 | encoding(rpc::StreamEncoding::IDENTITY) {} |
| 45 | WritableSinkImpl(kj::Exception reason) |
| 46 | : state(WritableSinkState::create<kj::Exception>(kj::mv(reason))), |
| 47 | encoding(rpc::StreamEncoding::IDENTITY) {} |
| 48 | |
| 49 | KJ_DISALLOW_COPY_AND_MOVE(WritableSinkImpl); |
| 50 | |
| 51 | virtual ~WritableSinkImpl() noexcept(false) { |
| 52 | if (!canceler.isEmpty()) { |
| 53 | canceler.cancel(KJ_EXCEPTION(DISCONNECTED, "stream was dropped")); |
| 54 | } |
| 55 | } |
| 56 | |
| 57 | kj::Promise<void> write(kj::ArrayPtr<const byte> buffer) override final { |
| 58 | throwIfErrored(); |
| 59 | KJ_IF_SOME(open, state.tryGetActiveUnsafe()) { |
| 60 | KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); |
| 61 | try { |
| 62 | co_return co_await canceler.wrap(encodeAndWrite(prepareWrite(kj::mv(open.stream)), buffer)); |
| 63 | } catch (...) { |
| 64 | handleOperationException(); |
| 65 | } |
| 66 | } |
| 67 | // Must be closed |
| 68 | JSG_FAIL_REQUIRE(Error, "Cannot write to a closed stream."); |
| 69 | } |
| 70 | |
| 71 | kj::Promise<void> write(kj::ArrayPtr<const kj::ArrayPtr<const byte>> pieces) override final { |
| 72 | throwIfErrored(); |
| 73 | KJ_IF_SOME(open, state.tryGetActiveUnsafe()) { |
| 74 | KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); |
| 75 | try { |
| 76 | co_return co_await canceler.wrap(encodeAndWrite(prepareWrite(kj::mv(open.stream)), pieces)); |
| 77 | } catch (...) { |
| 78 | handleOperationException(); |
| 79 | } |
| 80 | } |
| 81 | // Must be closed |
| 82 | JSG_FAIL_REQUIRE(Error, "Cannot write to a closed stream."); |
| 83 | } |
| 84 | |
| 85 | kj::Promise<void> end() override final { |
| 86 | throwIfErrored(); |
| 87 | if (state.is<Closed>()) { |
| 88 | co_return; |
| 89 | } |
| 90 | auto& open = state.requireActiveUnsafe(); |
| 91 | KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); |
| 92 | // The AsyncOutputStream interface does not yet have an end() method. |
| 93 | // Instead, we just drop it, signaling EOF. Eventually, it might get |
| 94 | // an end method, at which point we should use that instead. |
| 95 | try { |
| 96 | co_await canceler.wrap(endImpl(*open.stream)); |
| 97 | setClosed(); |
| 98 | co_return; |
| 99 | } catch (...) { |
| 100 | handleOperationException(); |
| 101 | } |
| 102 | } |
| 103 | |
| 104 | void abort(kj::Exception reason) override final { |
| 105 | canceler.cancel(reason.clone()); |
| 106 | setErrored(kj::mv(reason)); |
| 107 | } |
| 108 | |
| 109 | rpc::StreamEncoding disownEncodingResponsibility() override final { |
| 110 | auto prev = encoding; |
| 111 | encoding = rpc::StreamEncoding::IDENTITY; |
| 112 | return prev; |
| 113 | } |
| 114 | |
| 115 | rpc::StreamEncoding getEncoding() override final { |
| 116 | return encoding; |
| 117 | } |
| 118 | |
| 119 | protected: |
| 120 | // Throws the stored exception if in error state. |
| 121 | void throwIfErrored() { |
| 122 | KJ_IF_SOME(exception, state.tryGetErrorUnsafe()) { |
| 123 | kj::throwFatalException(exception.clone()); |
| 124 | } |
| 125 | } |
| 126 | |
| 127 | // Handles exceptions from write/end operations: stores the error and rethrows. |
| 128 | [[noreturn]] void handleOperationException() { |
| 129 | auto exception = kj::getCaughtExceptionAsKj(); |
| 130 | setErrored(exception.clone()); |
| 131 | kj::throwFatalException(kj::mv(exception)); |
| 132 | } |
| 133 | |
| 134 | virtual kj::AsyncOutputStream& prepareWrite(kj::Own<kj::AsyncOutputStream>&& inner) { |
| 135 | return setStream(kj::mv(inner)); |
| 136 | }; |
| 137 | |
| 138 | virtual kj::Promise<void> encodeAndWrite( |
| 139 | kj::AsyncOutputStream& output, kj::ArrayPtr<const kj::byte> data) { |
| 140 | co_await output.write(data); |
| 141 | } |
| 142 | |
| 143 | virtual kj::Promise<void> encodeAndWrite( |
| 144 | kj::AsyncOutputStream& output, kj::ArrayPtr<const kj::ArrayPtr<const kj::byte>> pieces) { |
| 145 | co_await output.write(pieces); |
| 146 | } |
| 147 | |
| 148 | virtual kj::Promise<void> endImpl(kj::AsyncOutputStream& output) { |
| 149 | // When using the default implementation, we assume IDENTITY encoding. |
| 150 | KJ_ASSERT(encoding == rpc::StreamEncoding::IDENTITY); |
| 151 | if (auto endable = dynamic_cast<EndableAsyncOutputStream*>(&output)) { |
| 152 | co_await endable->end(); |
| 153 | } else if (auto endable = dynamic_cast<capnp::ExplicitEndOutputStream*>(&output)) { |
| 154 | co_await endable->end(); |
| 155 | } |
| 156 | // By default there's nothing to flush. |
| 157 | co_return; |
| 158 | } |
| 159 | |
| 160 | void setClosed() { |
| 161 | state.transitionTo<Closed>(); |
| 162 | } |
| 163 | |
| 164 | void setErrored(kj::Exception&& ex) { |
| 165 | // Use forceTransitionTo because setErrored may be called when already |
| 166 | // in an error state (e.g., from write error handling). |
| 167 | state.forceTransitionTo<kj::Exception>(kj::mv(ex)); |
| 168 | } |
| 169 | |
| 170 | kj::AsyncOutputStream& setStream(kj::Own<kj::AsyncOutputStream> inner) { |
| 171 | auto& ret = *inner; |
| 172 | // Update the stream in place without a state transition. |
| 173 | // This is called from prepareWrite() which may wrap/transform the stream. |
| 174 | state.getUnsafe<Open>().stream = kj::mv(inner); |
| 175 | return ret; |
| 176 | } |
| 177 | |
| 178 | WritableSinkState& getState() { |
| 179 | return state; |
| 180 | } |
| 181 | |
| 182 | private: |
| 183 | WritableSinkState state; |
| 184 | rpc::StreamEncoding encoding; |
| 185 | kj::Canceler canceler; |
| 186 | }; |
| 187 | |
| 188 | // A wrapper around a native `kj::AsyncOutputStream` which knows the underlying encoding of the |
| 189 | // stream and optimizes pumps from `EncodedAsyncInputStream`. |
| 190 | // |
| 191 | // The inner will be held on to right up until either end() or abort() is called. |
| 192 | // This is important because some AsyncOutputStream implementations perform cleanup |
| 193 | // operations equivalent to end() in their destructors (for instance HttpChunkedEntityWriter). |
| 194 | // If we wait to clear the kj::Own when the EncodedAsyncOutputStream is destroyed, and the |
| 195 | // EncodedAsyncOutputStream is owned (for instance) by an IoOwn, then the lifetime of the |
| 196 | // inner may be extended past when it should. Eventually, kj::AsyncOutputStream should |
| 197 | // probably have a distinct end() method of its own that we can defer to, but until it |
| 198 | // does, it is important for us to release it as soon as end() or abort() are called. |
| 199 | class EncodedAsyncOutputStream final: public WritableSinkImpl { |
| 200 | public: |
| 201 | explicit EncodedAsyncOutputStream( |
| 202 | kj::Own<kj::AsyncOutputStream> inner, rpc::StreamEncoding encoding) |
| 203 | : WritableSinkImpl(kj::mv(inner), encoding) {} |
| 204 | |
| 205 | kj::Promise<void> endImpl(kj::AsyncOutputStream& output) override { |
| 206 | if (auto gzip = dynamic_cast<kj::GzipAsyncOutputStream*>(&output)) { |
| 207 | co_await gzip->end(); |
| 208 | } else if (auto br = dynamic_cast<kj::BrotliAsyncOutputStream*>(&output)) { |
| 209 | co_await br->end(); |
| 210 | } else if (auto endable = dynamic_cast<EndableAsyncOutputStream*>(&output)) { |
| 211 | co_await endable->end(); |
| 212 | } else if (auto endable = dynamic_cast<capnp::ExplicitEndOutputStream*>(&output)) { |
| 213 | co_await endable->end(); |
| 214 | } |
| 215 | // By default there's nothing to flush. |
| 216 | } |
| 217 | |
| 218 | kj::AsyncOutputStream& prepareWrite(kj::Own<kj::AsyncOutputStream>&& inner) override { |
| 219 | switch (disownEncodingResponsibility()) { |
| 220 | case rpc::StreamEncoding::GZIP: { |
| 221 | return setStream(kj::heap<kj::GzipAsyncOutputStream>(*inner).attach(kj::mv(inner))); |
| 222 | } |
| 223 | case rpc::StreamEncoding::BROTLI: { |
| 224 | return setStream(kj::heap<kj::BrotliAsyncOutputStream>(*inner).attach(kj::mv(inner))); |
| 225 | } |
| 226 | case rpc::StreamEncoding::IDENTITY: { |
| 227 | return setStream(kj::mv(inner)); |
| 228 | } |
| 229 | } |
| 230 | KJ_UNREACHABLE; |
| 231 | } |
| 232 | }; |
| 233 | |
| 234 | // A wrapper around a WritableSink that registers pending events with an IoContext. |
| 235 | class IoContextWritableSinkWrapper: public WritableSinkWrapper { |
| 236 | public: |
| 237 | IoContextWritableSinkWrapper(IoContext& ioContext, kj::Own<WritableSink> inner) |
| 238 | : WritableSinkWrapper(kj::mv(inner)), |
| 239 | ioContext(ioContext) {} |
| 240 | |
| 241 | kj::Promise<void> write(kj::ArrayPtr<const byte> buffer) override { |
| 242 | auto pending = ioContext.registerPendingEvent(); |
| 243 | KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { |
| 244 | co_await p; |
| 245 | } |
| 246 | co_await getInner().write(buffer); |
| 247 | } |
| 248 | |
| 249 | kj::Promise<void> write(kj::ArrayPtr<const kj::ArrayPtr<const byte>> pieces) override { |
| 250 | auto pending = ioContext.registerPendingEvent(); |
| 251 | KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { |
| 252 | co_await p; |
| 253 | } |
| 254 | co_await getInner().write(pieces); |
| 255 | } |
| 256 | |
| 257 | kj::Promise<void> end() override { |
| 258 | auto pending = ioContext.registerPendingEvent(); |
| 259 | KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { |
| 260 | co_await p; |
| 261 | } |
| 262 | co_await getInner().end(); |
| 263 | } |
| 264 | |
| 265 | private: |
| 266 | IoContext& ioContext; |
| 267 | }; |
| 268 | } // namespace |
| 269 | |
| 270 | kj::Own<WritableSink> newWritableSink(kj::Own<kj::AsyncOutputStream> inner) { |
| 271 | return kj::heap<WritableSinkImpl>(kj::mv(inner)); |
| 272 | } |
| 273 | |
| 274 | kj::Own<WritableSink> newClosedWritableSink() { |
| 275 | return kj::heap<WritableSinkImpl>(); |
| 276 | } |
| 277 | |
| 278 | kj::Own<WritableSink> newErroredWritableSink(kj::Exception reason) { |
| 279 | return kj::heap<WritableSinkImpl>(kj::mv(reason)); |
| 280 | } |
| 281 | |
| 282 | kj::Own<WritableSink> newNullWritableSink() { |
| 283 | return kj::heap<WritableSinkImpl>(newNullOutputStream()); |
| 284 | } |
| 285 | |
| 286 | kj::Own<WritableSink> newEncodedWritableSink( |
| 287 | rpc::StreamEncoding encoding, kj::Own<kj::AsyncOutputStream> inner) { |
| 288 | return kj::heap<EncodedAsyncOutputStream>(kj::mv(inner), encoding); |
| 289 | } |
| 290 | |
| 291 | kj::Own<WritableSink> newIoContextWrappedWritableSink( |
| 292 | IoContext& ioContext, kj::Own<WritableSink> inner) { |
| 293 | return kj::heap<IoContextWritableSinkWrapper>(ioContext, kj::mv(inner)); |
| 294 | } |
| 295 | |
| 296 | } // namespace workerd::api::streams |