Skip to content
File

Blob: src/workerd/util/stream-utils.c++

8.4 KB
1#include "stream-utils.h"
2 
3#include <kj/common.h>
4#include <kj/debug.h>
5#include <kj/exception.h>
6#include <kj/one-of.h>
7 
8namespace workerd {
9 
10namespace {
11 
12// An AsyncInputStream implementation that reads from an in-memory buffer.
13// This is optimized for the case where the entire contents are available
14// up-front, so it doesn't do any dynamic memory allocation or copying.
15// It also supports optimized teeing when the backing storage is provided.
16class MemoryInputStream final: public kj::AsyncInputStream {
17 private:
18 struct OwnedBacking: public kj::Refcounted {
19 kj::Own<void> backing;
20 OwnedBacking(kj::Own<void>&& backing): backing(kj::mv(backing)) {}
21 };
22 
23 public:
24 MemoryInputStream(
25 kj::ArrayPtr<const kj::byte> data, kj::Maybe<kj::Own<void>> maybeBacking = kj::none)
26 : data(data),
27 // Note that we don't actually check that maybeBacking actually owns the
28 // memory that `data` points to. It is the caller's responsibility to ensure
29 // this is the case if they want teeing to be safely supported.
30 ownedBacking(maybeBacking.map(
31 [](kj::Own<void>& backing) mutable { return kj::rc<OwnedBacking>(kj::mv(backing)); })) {
32 }
33 MemoryInputStream(kj::ArrayPtr<const kj::byte> data, kj::Rc<OwnedBacking> ownedBacking)
34 : data(data),
35 ownedBacking(kj::mv(ownedBacking)) {}
36 
37 kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override {
38 auto ptr = kj::arrayPtr<kj::byte>(static_cast<kj::byte*>(buffer), maxBytes);
39 size_t toRead = kj::min(data.size(), ptr.size());
40 if (toRead == 0) return toRead;
41 ptr.first(toRead).copyFrom(data.first(toRead));
42 data = data.slice(toRead);
43 return toRead;
44 }
45 
46 kj::Maybe<uint64_t> tryGetLength() override {
47 return data.size();
48 }
49 
50 kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override {
51 // An optimized pumpTo... we know we have all the data right here. We can
52 // just write it all at once up to `amount`.
53 uint64_t toRead = kj::min(data.size(), amount);
54 if (toRead == 0) {
55 co_return toRead;
56 }
57 co_await output.write(data.first(toRead));
58 data = data.slice(toRead);
59 co_return toRead;
60 }
61 
62 kj::Maybe<kj::Own<AsyncInputStream>> tryTee(uint64_t limit = kj::maxValue) override {
63 // If a MemoryInputStream is holding onto backing storage, then we can safely
64 // tee it here, allowing us to avoid the default tee implementation which needs
65 // additional buffering. Tee'ing just becomes a matter of sharing the backing
66 // storage and the data slice directly. This allows us to avoid any additional
67 // buffering in unread stream branches since all of the data is already in memory
68 // anyway. If we're not holding onto the backing storage, then we cannot safely
69 // assume that the a tee branch will safely be able to read the data, so we'll
70 // fall back to the default kj::newTee implementation.
71 KJ_IF_SOME(owned, ownedBacking) {
72 return kj::heap<MemoryInputStream>(data, owned.addRef());
73 }
74 return kj::none;
75 }
76 
77 private:
78 kj::ArrayPtr<const kj::byte> data;
79 kj::Maybe<kj::Rc<OwnedBacking>> ownedBacking;
80};
81 
82class NeuterableInputStreamImpl final: public NeuterableInputStream {
83 public:
84 NeuterableInputStreamImpl(kj::AsyncInputStream& inner): inner(&inner) {}
85 
86 void neuter(kj::Exception exception) override {
87 if (inner.is<kj::AsyncInputStream*>()) {
88 inner = exception.clone();
89 if (!canceler.isEmpty()) {
90 canceler.cancel(kj::mv(exception));
91 }
92 }
93 }
94 
95 kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override {
96 return canceler.wrap(getStream().tryRead(buffer, minBytes, maxBytes));
97 }
98 kj::Maybe<uint64_t> tryGetLength() override {
99 return getStream().tryGetLength();
100 }
101 kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override {
102 return canceler.wrap(getStream().pumpTo(output, amount));
103 }
104 
105 private:
106 kj::OneOf<kj::AsyncInputStream*, kj::Exception> inner;
107 kj::Canceler canceler;
108 
109 kj::AsyncInputStream& getStream() {
110 KJ_SWITCH_ONEOF(inner) {
111 KJ_CASE_ONEOF(stream, kj::AsyncInputStream*) {
112 return *stream;
113 }
114 KJ_CASE_ONEOF(exception, kj::Exception) {
115 kj::throwFatalException(exception.clone());
116 }
117 }
118 KJ_UNREACHABLE;
119 }
120};
121 
122class NeuterableIoStreamImpl final: public NeuterableIoStream {
123 public:
124 NeuterableIoStreamImpl(kj::AsyncIoStream& inner): inner(&inner) {}
125 
126 void neuter(kj::Exception reason) override {
127 if (inner.is<kj::AsyncIoStream*>()) {
128 inner = reason.clone();
129 if (!canceler.isEmpty()) {
130 canceler.cancel(kj::mv(reason));
131 }
132 }
133 }
134 
135 // AsyncInputStream
136 
137 kj::Promise<size_t> tryRead(void* buffer, size_t minBytes, size_t maxBytes) override {
138 return canceler.wrap(getStream().tryRead(buffer, minBytes, maxBytes));
139 }
140 kj::Maybe<uint64_t> tryGetLength() override {
141 return getStream().tryGetLength();
142 }
143 kj::Promise<uint64_t> pumpTo(kj::AsyncOutputStream& output, uint64_t amount) override {
144 return canceler.wrap(getStream().pumpTo(output, amount));
145 }
146 
147 // AsyncOutputStream
148 
149 kj::Promise<void> write(kj::ArrayPtr<const kj::byte> buffer) override {
150 return canceler.wrap(getStream().write(buffer));
151 }
152 kj::Promise<void> write(kj::ArrayPtr<const kj::ArrayPtr<const kj::byte>> pieces) override {
153 return canceler.wrap(getStream().write(pieces));
154 }
155 kj::Maybe<kj::Promise<uint64_t>> tryPumpFrom(
156 kj::AsyncInputStream& input, uint64_t amount) override {
157 return getStream().tryPumpFrom(input, amount).map([this](kj::Promise<uint64_t> promise) {
158 return canceler.wrap(kj::mv(promise));
159 });
160 }
161 kj::Promise<void> whenWriteDisconnected() override {
162 return canceler.wrap(getStream().whenWriteDisconnected());
163 }
164 
165 // AsyncIoStream
166 
167 void shutdownWrite() override {
168 getStream().shutdownWrite();
169 };
170 void abortRead() override {
171 getStream().abortRead();
172 }
173 void getsockopt(int level, int option, void* value, kj::uint* length) override {
174 getStream().getsockopt(level, option, value, length);
175 }
176 void setsockopt(int level, int option, const void* value, kj::uint length) override {
177 getStream().setsockopt(level, option, value, length);
178 }
179 void getsockname(struct sockaddr* addr, kj::uint* length) override {
180 getStream().getsockname(addr, length);
181 }
182 void getpeername(struct sockaddr* addr, kj::uint* length) override {
183 getStream().getpeername(addr, length);
184 }
185 virtual kj::Maybe<int> getFd() const override {
186 return getStream().getFd();
187 }
188 
189 private:
190 kj::OneOf<kj::AsyncIoStream*, kj::Exception> inner;
191 kj::Canceler canceler;
192 
193 kj::AsyncIoStream& getStream() {
194 KJ_IF_SOME(stream, inner.tryGet<kj::AsyncIoStream*>()) {
195 return *stream;
196 }
197 kj::throwFatalException(inner.get<kj::Exception>().clone());
198 }
199 kj::AsyncIoStream& getStream() const {
200 KJ_IF_SOME(stream, inner.tryGet<kj::AsyncIoStream*>()) {
201 return *stream;
202 }
203 kj::throwFatalException(inner.get<kj::Exception>().clone());
204 }
205};
206 
207// The kj::NullStream instance is stateless, discards all writes, and returns
208// EOF on all reads. We can, therefore, safely share a single static global
209// instance instead of allocating a new one each time.
210static kj::NullStream nullStream{};
211 
212} // namespace
213 
214kj::AsyncOutputStream& getGlobalNullOutputStream() {
215 return nullStream;
216}
217 
218kj::Own<kj::AsyncIoStream> newNullIoStream() {
219 return kj::Own<kj::AsyncIoStream>(&nullStream, kj::NullDisposer::instance);
220}
221 
222kj::Own<kj::AsyncInputStream> newNullInputStream() {
223 return kj::Own<kj::AsyncInputStream>(&nullStream, kj::NullDisposer::instance);
224}
225 
226kj::Own<kj::AsyncOutputStream> newNullOutputStream() {
227 return kj::Own<kj::AsyncOutputStream>(&nullStream, kj::NullDisposer::instance);
228}
229 
230kj::Own<kj::AsyncInputStream> newMemoryInputStream(
231 kj::ArrayPtr<const kj::byte> data, kj::Maybe<kj::Own<void>> maybeBacking) {
232 return kj::heap<MemoryInputStream>(data, kj::mv(maybeBacking));
233}
234 
235kj::Own<kj::AsyncInputStream> newMemoryInputStream(
236 kj::StringPtr data, kj::Maybe<kj::Own<void>> maybeBacking) {
237 return kj::heap<MemoryInputStream>(data.asBytes(), kj::mv(maybeBacking));
238}
239 
240kj::Own<NeuterableInputStream> newNeuterableInputStream(kj::AsyncInputStream& inner) {
241 return kj::refcounted<NeuterableInputStreamImpl>(inner);
242}
243 
244kj::Own<NeuterableIoStream> newNeuterableIoStream(kj::AsyncIoStream& inner) {
245 return kj::heap<NeuterableIoStreamImpl>(inner);
246}
247 
248} // namespace workerd