Skip to content
File

Blob: src/workerd/api/streams/writable-sink.c++

10.0 KB
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 
12namespace workerd::api::streams {
13 
14namespace {
15struct Closed {
16 static constexpr kj::StringPtr NAME KJ_UNUSED = "closed"_kj;
17};
18 
19struct 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.
28using 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.
36class 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.
199class 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.
235class 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 
270kj::Own<WritableSink> newWritableSink(kj::Own<kj::AsyncOutputStream> inner) {
271 return kj::heap<WritableSinkImpl>(kj::mv(inner));
272}
273 
274kj::Own<WritableSink> newClosedWritableSink() {
275 return kj::heap<WritableSinkImpl>();
276}
277 
278kj::Own<WritableSink> newErroredWritableSink(kj::Exception reason) {
279 return kj::heap<WritableSinkImpl>(kj::mv(reason));
280}
281 
282kj::Own<WritableSink> newNullWritableSink() {
283 return kj::heap<WritableSinkImpl>(newNullOutputStream());
284}
285 
286kj::Own<WritableSink> newEncodedWritableSink(
287 rpc::StreamEncoding encoding, kj::Own<kj::AsyncOutputStream> inner) {
288 return kj::heap<EncodedAsyncOutputStream>(kj::mv(inner), encoding);
289}
290 
291kj::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