#include "writable-sink.h" #include #include #include #include #include #include #include namespace workerd::api::streams { namespace { struct Closed { static constexpr kj::StringPtr NAME KJ_UNUSED = "closed"_kj; }; struct Open { static constexpr kj::StringPtr NAME KJ_UNUSED = "open"_kj; kj::Own stream; }; // State machine for tracking writable sink lifecycle: // Open -> Closed (normal close via end()) // Open -> kj::Exception (error via abort() or write failure) // Closed is terminal, kj::Exception is implicitly terminal via ErrorState. using WritableSinkState = StateMachine, ErrorState, ActiveState, Open, Closed, kj::Exception>; // The base implementation of WritableSink. This is not exposed publicly. class WritableSinkImpl: public WritableSink { public: WritableSinkImpl(kj::Own inner, rpc::StreamEncoding encoding = rpc::StreamEncoding::IDENTITY) : state(WritableSinkState::create(kj::mv(inner))), encoding(encoding) {} WritableSinkImpl() : state(WritableSinkState::create()), encoding(rpc::StreamEncoding::IDENTITY) {} WritableSinkImpl(kj::Exception reason) : state(WritableSinkState::create(kj::mv(reason))), encoding(rpc::StreamEncoding::IDENTITY) {} KJ_DISALLOW_COPY_AND_MOVE(WritableSinkImpl); virtual ~WritableSinkImpl() noexcept(false) { if (!canceler.isEmpty()) { canceler.cancel(KJ_EXCEPTION(DISCONNECTED, "stream was dropped")); } } kj::Promise write(kj::ArrayPtr buffer) override final { throwIfErrored(); KJ_IF_SOME(open, state.tryGetActiveUnsafe()) { KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); try { co_return co_await canceler.wrap(encodeAndWrite(prepareWrite(kj::mv(open.stream)), buffer)); } catch (...) { handleOperationException(); } } // Must be closed JSG_FAIL_REQUIRE(Error, "Cannot write to a closed stream."); } kj::Promise write(kj::ArrayPtr> pieces) override final { throwIfErrored(); KJ_IF_SOME(open, state.tryGetActiveUnsafe()) { KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); try { co_return co_await canceler.wrap(encodeAndWrite(prepareWrite(kj::mv(open.stream)), pieces)); } catch (...) { handleOperationException(); } } // Must be closed JSG_FAIL_REQUIRE(Error, "Cannot write to a closed stream."); } kj::Promise end() override final { throwIfErrored(); if (state.is()) { co_return; } auto& open = state.requireActiveUnsafe(); KJ_REQUIRE(canceler.isEmpty(), "jsg.Error: Stream is already being written to"); // The AsyncOutputStream interface does not yet have an end() method. // Instead, we just drop it, signaling EOF. Eventually, it might get // an end method, at which point we should use that instead. try { co_await canceler.wrap(endImpl(*open.stream)); setClosed(); co_return; } catch (...) { handleOperationException(); } } void abort(kj::Exception reason) override final { canceler.cancel(reason.clone()); setErrored(kj::mv(reason)); } rpc::StreamEncoding disownEncodingResponsibility() override final { auto prev = encoding; encoding = rpc::StreamEncoding::IDENTITY; return prev; } rpc::StreamEncoding getEncoding() override final { return encoding; } protected: // Throws the stored exception if in error state. void throwIfErrored() { KJ_IF_SOME(exception, state.tryGetErrorUnsafe()) { kj::throwFatalException(exception.clone()); } } // Handles exceptions from write/end operations: stores the error and rethrows. [[noreturn]] void handleOperationException() { auto exception = kj::getCaughtExceptionAsKj(); setErrored(exception.clone()); kj::throwFatalException(kj::mv(exception)); } virtual kj::AsyncOutputStream& prepareWrite(kj::Own&& inner) { return setStream(kj::mv(inner)); }; virtual kj::Promise encodeAndWrite( kj::AsyncOutputStream& output, kj::ArrayPtr data) { co_await output.write(data); } virtual kj::Promise encodeAndWrite( kj::AsyncOutputStream& output, kj::ArrayPtr> pieces) { co_await output.write(pieces); } virtual kj::Promise endImpl(kj::AsyncOutputStream& output) { // When using the default implementation, we assume IDENTITY encoding. KJ_ASSERT(encoding == rpc::StreamEncoding::IDENTITY); if (auto endable = dynamic_cast(&output)) { co_await endable->end(); } else if (auto endable = dynamic_cast(&output)) { co_await endable->end(); } // By default there's nothing to flush. co_return; } void setClosed() { state.transitionTo(); } void setErrored(kj::Exception&& ex) { // Use forceTransitionTo because setErrored may be called when already // in an error state (e.g., from write error handling). state.forceTransitionTo(kj::mv(ex)); } kj::AsyncOutputStream& setStream(kj::Own inner) { auto& ret = *inner; // Update the stream in place without a state transition. // This is called from prepareWrite() which may wrap/transform the stream. state.getUnsafe().stream = kj::mv(inner); return ret; } WritableSinkState& getState() { return state; } private: WritableSinkState state; rpc::StreamEncoding encoding; kj::Canceler canceler; }; // A wrapper around a native `kj::AsyncOutputStream` which knows the underlying encoding of the // stream and optimizes pumps from `EncodedAsyncInputStream`. // // The inner will be held on to right up until either end() or abort() is called. // This is important because some AsyncOutputStream implementations perform cleanup // operations equivalent to end() in their destructors (for instance HttpChunkedEntityWriter). // If we wait to clear the kj::Own when the EncodedAsyncOutputStream is destroyed, and the // EncodedAsyncOutputStream is owned (for instance) by an IoOwn, then the lifetime of the // inner may be extended past when it should. Eventually, kj::AsyncOutputStream should // probably have a distinct end() method of its own that we can defer to, but until it // does, it is important for us to release it as soon as end() or abort() are called. class EncodedAsyncOutputStream final: public WritableSinkImpl { public: explicit EncodedAsyncOutputStream( kj::Own inner, rpc::StreamEncoding encoding) : WritableSinkImpl(kj::mv(inner), encoding) {} kj::Promise endImpl(kj::AsyncOutputStream& output) override { if (auto gzip = dynamic_cast(&output)) { co_await gzip->end(); } else if (auto br = dynamic_cast(&output)) { co_await br->end(); } else if (auto endable = dynamic_cast(&output)) { co_await endable->end(); } else if (auto endable = dynamic_cast(&output)) { co_await endable->end(); } // By default there's nothing to flush. } kj::AsyncOutputStream& prepareWrite(kj::Own&& inner) override { switch (disownEncodingResponsibility()) { case rpc::StreamEncoding::GZIP: { return setStream(kj::heap(*inner).attach(kj::mv(inner))); } case rpc::StreamEncoding::BROTLI: { return setStream(kj::heap(*inner).attach(kj::mv(inner))); } case rpc::StreamEncoding::IDENTITY: { return setStream(kj::mv(inner)); } } KJ_UNREACHABLE; } }; // A wrapper around a WritableSink that registers pending events with an IoContext. class IoContextWritableSinkWrapper: public WritableSinkWrapper { public: IoContextWritableSinkWrapper(IoContext& ioContext, kj::Own inner) : WritableSinkWrapper(kj::mv(inner)), ioContext(ioContext) {} kj::Promise write(kj::ArrayPtr buffer) override { auto pending = ioContext.registerPendingEvent(); KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { co_await p; } co_await getInner().write(buffer); } kj::Promise write(kj::ArrayPtr> pieces) override { auto pending = ioContext.registerPendingEvent(); KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { co_await p; } co_await getInner().write(pieces); } kj::Promise end() override { auto pending = ioContext.registerPendingEvent(); KJ_IF_SOME(p, ioContext.waitForOutputLocksIfNecessary()) { co_await p; } co_await getInner().end(); } private: IoContext& ioContext; }; } // namespace kj::Own newWritableSink(kj::Own inner) { return kj::heap(kj::mv(inner)); } kj::Own newClosedWritableSink() { return kj::heap(); } kj::Own newErroredWritableSink(kj::Exception reason) { return kj::heap(kj::mv(reason)); } kj::Own newNullWritableSink() { return kj::heap(newNullOutputStream()); } kj::Own newEncodedWritableSink( rpc::StreamEncoding encoding, kj::Own inner) { return kj::heap(kj::mv(inner), encoding); } kj::Own newIoContextWrappedWritableSink( IoContext& ioContext, kj::Own inner) { return kj::heap(ioContext, kj::mv(inner)); } } // namespace workerd::api::streams